From 0d7fc34381e46007bea9f8c622704aee80a83050 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Fri, 4 Sep 2026 08:03:40 +0200 Subject: [PATCH 01/18] test(admin): add red tests for the upstream provider report endpoint Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01QS9ws7DWUM9ryiroWSWwo1 --- .../test_admin_upstream_provider_report.py | 603 ++++++++++++++++++ 1 file changed, 603 insertions(+) create mode 100644 tests/integration/test_admin_upstream_provider_report.py diff --git a/tests/integration/test_admin_upstream_provider_report.py b/tests/integration/test_admin_upstream_provider_report.py new file mode 100644 index 00000000..c64fcaba --- /dev/null +++ b/tests/integration/test_admin_upstream_provider_report.py @@ -0,0 +1,603 @@ +"""Certification report for a configured upstream provider. + +Covers ``GET /admin/api/upstream-providers/{provider_id}/report``: the row +contract shape, and the four pricing rows it carries — +``pricing.served_matches_configured``, ``pricing.sats_pricing_present``, +``pricing.enabled_models_served`` and ``pricing.cache_rate``. Each row is +computed from the DB row plus the in-process served map; none of them make a +network call. +""" + +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone +from typing import Any +from unittest.mock import patch + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + +PRICING_ROW_IDS = ( + "pricing.served_matches_configured", + "pricing.sats_pricing_present", + "pricing.enabled_models_served", + "pricing.cache_rate", +) + +ARCHITECTURE = { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, +} + + +def _admin_headers() -> dict[str, str]: + token = "test-admin-upstream-report-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +async def _make_provider( + session: AsyncSession, + *, + slug: str | None = None, + provider_fee: float = 1.0, + base_url: str = "https://report-upstream.example/v1", + api_key: str = "test-key", +) -> UpstreamProviderRow: + provider = UpstreamProviderRow( + provider_type="generic", + base_url=base_url, + api_key=api_key, + provider_fee=provider_fee, + slug=slug, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + assert provider.id is not None + return provider + + +def _pricing(**overrides: object) -> dict[str, object]: + pricing: dict[str, object] = { + "prompt": 1.4e-7, + "completion": 2.8e-7, + "request": 0.0, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, + "input_cache_read": 0.0, + "input_cache_write": 0.0, + } + pricing.update(overrides) + return pricing + + +def _model_row( + provider_id: int, + *, + model_id: str, + pricing: dict[str, object], + enabled: bool = True, +) -> ModelRow: + return ModelRow( + id=model_id, + name=model_id, + description="d", + created=0, + context_length=8192, + architecture=json.dumps(ARCHITECTURE), + pricing=json.dumps(pricing), + upstream_provider_id=provider_id, + enabled=enabled, + # A self-alias, same as the admin write edge stores by default — + # ``get_effective_forwarded_model_id`` treats this as "no distinct + # forwarded id" so it does not register a second routable alias. + forwarded_model_id=model_id, + ) + + +def _row_ids(rows: list[dict[str, Any]]) -> list[str]: + return [row["id"] for row in rows] + + +def _find_row(rows: list[dict[str, Any]], row_id: str) -> dict[str, Any]: + for row in rows: + if row["id"] == row_id: + return row + raise AssertionError(f"row {row_id!r} not found in {_row_ids(rows)!r}") + + +def _pid(provider: UpstreamProviderRow) -> int: + """Narrow a persisted row's optional primary key for typed call sites.""" + assert provider.id is not None + return provider.id + + +async def _get_report(client: AsyncClient, provider_ref: str | int) -> Any: + return await client.get( + f"/admin/api/upstream-providers/{provider_ref}/report", + headers=_admin_headers(), + ) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_report_requires_admin_auth( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """Sanity check: the report sits behind the same gate as the rest of + ``core/admin.py``. This already passes against the not-implemented stub + because ``require_admin_api`` runs as a dependency before the route body + — it is included for completeness, not as a red proof. + """ + provider = await _make_provider(integration_session) + + resp = await integration_client.get( + f"/admin/api/upstream-providers/{_pid(provider)}/report" + ) + + assert resp.status_code == 403 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_report_unknown_provider_returns_404( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + resp = await _get_report(integration_client, 999_999_999) + + assert resp.status_code == 404 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_report_row_contract_shape_and_order( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """The row contract from the report-contract spec: stable top-level keys, + a fixed row order with the four pricing rows first, and every row + carrying id/status/title/detail/evidence with status in {ok, warn, fail}. + """ + provider = await _make_provider(integration_session, slug="report-shape-provider") + integration_session.add( + _model_row(_pid(provider), model_id="shape-model", pricing=_pricing()) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + assert provider.slug is not None + resp = await _get_report(integration_client, provider.slug) + + assert resp.status_code == 200, resp.text + body = resp.json() + + # The numeric id, not an echo of whatever ref (slug, here) the request + # used to look the provider up — matches ``_serialize_provider``'s "id". + assert body["provider_id"] == provider.id + generated_at = body["generated_at"] + # Must parse as an ISO-8601 timestamp; a trailing "Z" is not accepted by + # ``fromisoformat`` on its own. + parsed_generated_at = datetime.fromisoformat(generated_at.replace("Z", "+00:00")) + # Freshly generated, not a stale cached/hardcoded value. + assert abs((datetime.now(timezone.utc) - parsed_generated_at).total_seconds()) < 60 + + rows = body["rows"] + assert _row_ids(rows)[:4] == list(PRICING_ROW_IDS) + for row in rows: + assert set(row) >= {"id", "status", "title", "detail", "evidence"} + assert row["status"] in {"ok", "warn", "fail"} + assert isinstance(row["title"], str) and row["title"] + assert isinstance(row["detail"], str) and row["detail"] + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_served_matches_configured_ok_when_prices_agree( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider = await _make_provider(integration_session, provider_fee=1.05) + integration_session.add( + _model_row( + _pid(provider), + model_id="agree-model", + pricing=_pricing(prompt=2e-7, completion=4e-7), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.served_matches_configured") + assert row["status"] == "ok", row + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_served_matches_configured_zero_vs_zero_is_ok( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A price of zero on both sides is agreement, not a legitimacy check.""" + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id="free-model", + pricing=_pricing(prompt=0.0, completion=0.0), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.served_matches_configured") + assert row["status"] == "ok", row + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_served_matches_configured_fails_on_stale_served_map( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """Drift the DB row without refreshing the served map — the same shape + of staleness a writer that bypasses ``core/admin.py`` would leave behind. + "Configured" (built fresh from the row) must then disagree with "served" + (built earlier, still in-process) with no epsilon. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id="drift-model", + pricing=_pricing(prompt=1e-7, completion=2e-7), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + stored = await integration_session.get(ModelRow, ("drift-model", _pid(provider))) + assert stored is not None + stored.pricing = json.dumps(_pricing(prompt=9e-7, completion=2e-7)) + integration_session.add(stored) + await integration_session.commit() + # Deliberately no reinitialize_upstreams() here: the served map must stay + # stale for this to be a meaningful drift case. + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.served_matches_configured") + assert row["status"] == "fail", row + assert row["evidence"] is not None + assert "drift-model" in json.dumps(row["evidence"]) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_sats_pricing_present_ok_when_conversion_succeeds( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider = await _make_provider(integration_session) + integration_session.add( + _model_row(_pid(provider), model_id="sats-ok-model", pricing=_pricing()) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.sats_pricing_present") + assert row["status"] == "ok", row + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_sats_pricing_present_fails_when_btc_feed_is_swallowed( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """``_update_model_sats_pricing`` swallows every exception and leaves the + served model with ``sats_pricing=None``. This must surface here rather + than silently advertising models with no sats price. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row(_pid(provider), model_id="sats-fail-model", pricing=_pricing()) + ) + await integration_session.commit() + with patch( + "routstr.payment.models.sats_usd_price", + side_effect=RuntimeError("btc feed unavailable"), + ): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.sats_pricing_present") + assert row["status"] == "fail", row + assert "sats-fail-model" in json.dumps(row["evidence"]) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_enabled_models_served_ok_when_all_enabled_models_are_served( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider = await _make_provider(integration_session) + integration_session.add( + _model_row(_pid(provider), model_id="served-model", pricing=_pricing()) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.enabled_models_served") + assert row["status"] == "ok", row + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_enabled_models_served_fails_when_enabled_model_has_unusable_pricing( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A negative rate makes ``has_usable_pricing`` false, so the algorithm + withholds the model from the served map even though the DB row is + enabled — exactly the "enabled but never served" case this row exists + to catch, and it must not require an upstream that stopped listing the + model to reproduce. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id="unusable-price-model", + pricing=_pricing(prompt=-1.0), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.enabled_models_served") + assert row["status"] == "fail", row + assert "unusable-price-model" in json.dumps(row["evidence"]) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_report_survives_a_model_row_that_fails_to_parse( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """ "A row never throws": a stored row even malformed enough that + ``_build_model_from_row`` raises on it (bad JSON, in this case — the same + shape of corruption a legacy writer can leave) must become a ``fail`` row + with the exception described, not a 500 that takes out the whole report. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row(_pid(provider), model_id="good-model", pricing=_pricing()) + ) + broken = _model_row(_pid(provider), model_id="broken-model", pricing=_pricing()) + broken.pricing = "{not valid json" + integration_session.add(broken) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + body = resp.json() + assert _row_ids(body["rows"])[:4] == list(PRICING_ROW_IDS) + row = _find_row(body["rows"], "pricing.served_matches_configured") + assert row["status"] == "fail", row + assert "broken-model" in json.dumps(row["evidence"]) + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_cache_rate_ignores_an_enabled_model_that_is_not_served( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A negative price holds a model back from the served map even though + its row is enabled (see ``test_enabled_models_served_fails_when_...``). + ``pricing.cache_rate`` must not certify a cache rate for a model that + isn't actually being served — it should skip it, not count it, and + certainly not report ``ok`` for a model nothing will ever bill through. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id="unserved-model", + pricing=_pricing(prompt=-1.0), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.cache_rate") + assert row["status"] == "ok", row + assert row["evidence"]["checked"] == 0 + assert "unserved-model" not in json.dumps(row["evidence"]) + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("model_id", ["gpt-4o", "deepseek-chat"]) +async def test_cache_rate_warns_when_backfill_only_supplies_the_read_rate( + integration_client: AsyncClient, + integration_session: AsyncSession, + model_id: str, +) -> None: + """The row is computed from ``backfill_cache_pricing(row.id, pricing)`` at + serve time, not from the raw DB row. Both ``gpt-4o`` and DeepSeek chat + models are stored with ``input_cache_read=0`` (the OpenRouter feed omits + it) and litellm's cost map fills that in — reading the raw row instead + would falsely flag the read rate as unknown, which is the defect this + row's spec was corrected to avoid. + + litellm's cost map has no ``cache_creation_input_token_cost`` entry for + either model, so the write rate stays unbackfilled: the row must still + ``warn`` (a real, if partial, gap) rather than call this ``ok``. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id=model_id, + pricing=_pricing(prompt=2.5e-6, completion=1e-5, input_cache_read=0.0), + ) + ) + await integration_session.commit() + + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.cache_rate") + assert row["status"] == "warn", row + evidence_text = json.dumps(row["evidence"]) + assert model_id in evidence_text + assert "input_cache_write" in evidence_text + assert "input_cache_read" not in evidence_text + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_cache_rate_ok_when_backfill_supplies_both_rates( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """``claude-sonnet-4-5`` has both a cache-read and a cache-creation + (write) rate in litellm's cost map, so once both are backfilled the row + must be ``ok`` — this is the counterpart to the partial-coverage case + above, proving ``ok`` is reachable and not just a status the row never + returns once both rates are checked. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id="claude-sonnet-4-5", + pricing=_pricing(prompt=3e-6, completion=1.5e-5, input_cache_read=0.0), + ) + ) + await integration_session.commit() + + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.cache_rate") + assert row["status"] == "ok", row + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_cache_rate_warns_when_rate_missing_and_unknown_to_litellm( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """No cache rate, and litellm has never heard of the model: the report + has no persisted probe result yet (that lands with the cost probe), so + this must be ``warn``, never ``fail`` — ``fail`` needs the probe to know + the upstream is token-billed. + """ + provider = await _make_provider(integration_session) + integration_session.add( + _model_row( + _pid(provider), + model_id="totally-custom-self-hosted-model", + pricing=_pricing(), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.cache_rate") + assert row["status"] == "warn", row + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_report_rows_are_scoped_to_the_requested_provider( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """A second provider's broken model must not leak into this provider's + aggregate row — each row is scoped to the provider named in the URL. + """ + provider_a = await _make_provider(integration_session, slug="scope-provider-a") + provider_b = await _make_provider( + integration_session, + slug="scope-provider-b", + base_url="https://report-upstream-b.example/v1", + api_key="test-key-b", + ) + + integration_session.add( + _model_row( + _pid(provider_a), + model_id="scope-a-model", + pricing=_pricing(prompt=1e-7, completion=2e-7), + ) + ) + integration_session.add( + _model_row( + _pid(provider_b), + model_id="scope-b-model", + pricing=_pricing(prompt=-1.0), + ) + ) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await _get_report(integration_client, _pid(provider_a)) + + assert resp.status_code == 200, resp.text + row = _find_row(resp.json()["rows"], "pricing.enabled_models_served") + assert row["status"] == "ok", row + # Evidence must actually be inspectable here, not merely absent — an "ok" + # row that reports ``evidence: None`` would make the leak check below + # vacuously true (``"x" not in json.dumps(None)`` is always True) instead + # of proving provider_b's model never entered provider_a's row. + assert row["evidence"] is not None + assert "scope-b-model" not in json.dumps(row["evidence"]) From f6d0d64345cc9af2a428aeb3dc2a037174e9ae76 Mon Sep 17 00:00:00 2001 From: Jeroen Ubbink Date: Fri, 4 Sep 2026 08:03:47 +0200 Subject: [PATCH 02/18] feat(admin): add the upstream provider certification report endpoint Adds GET /admin/api/upstream-providers/{provider_id}/report, returning a fixed-order set of rows about a configured provider's pricing health: whether the served price matches the configured one, whether a sats price was actually computed for each served model, whether every enabled model is being served, and whether both a cache-read and a cache-write rate are known for it (falling back to litellm's bundled cost map the same way the serve path does). Every row is derived from the database row plus the in-process served map, computed once per model and shared across all four rows, so the endpoint never makes a network call, never spends, and never re-derives the same fact four times over. A row that hits a malformed stored value reports itself as a fail row with the error instead of taking the whole request down; the cache-rate row is scoped to models that are actually being served, matching its three sibling rows. Co-Authored-By: Claude Sonnet 5 Claude-Session: https://claude.ai/code/session_01QS9ws7DWUM9ryiroWSWwo1 --- routstr/core/admin.py | 256 +++++++++++++++++++++++++++++++++++++++++- 1 file changed, 255 insertions(+), 1 deletion(-) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index cc140f6c..290480ac 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,6 +1,8 @@ import json import re import secrets +from collections.abc import Sequence +from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path @@ -12,11 +14,13 @@ from sqlmodel.ext.asyncio.session import AsyncSession from ..payment.models import ( REQUIRED_PRICING_FIELDS, + Model, + _build_model_from_row, _row_to_model, list_models, ) from ..payment.rates import BILLABLE_PRICING_FIELDS, coerce_rate -from ..proxy import refresh_model_maps, reinitialize_upstreams +from ..proxy import get_candidates, refresh_model_maps, reinitialize_upstreams from ..wallet import fetch_all_balances, send_token, token_mint_url from . import vault from .db import ( @@ -1246,6 +1250,256 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: } +def _served_model_for_provider(model_id: str, provider_pk: int) -> Model | None: + """The model this provider serves for ``model_id``, or ``None``. + + ``get_candidates`` returns every provider's candidate for the alias (a + model id can be served by more than one configured provider); narrow to + the one this report is about. + """ + for model, _upstream in get_candidates(model_id) or []: + if model.upstream_provider_id == provider_pk: + return model + return None + + +@dataclass +class _ModelEvaluation: + """One enabled model row's facts, built once and shared by every row. + + ``configured`` is the fee-applied USD view built fresh from the row, or + ``None`` when the stored row could not be parsed (``build_error`` then + carries the exception). ``served`` is this provider's live candidate for + the model, or ``None`` when it is not being served at all — e.g. an + unusable stored price holds it back from the served map even though the + row itself is enabled. + """ + + model_id: str + configured: Model | None + build_error: str | None + served: Model | None + + +def _evaluate_model_row( + row: ModelRow, provider: UpstreamProviderRow, provider_pk: int +) -> _ModelEvaluation: + try: + configured: Model | None = _build_model_from_row( + row, apply_provider_fee=True, provider_fee=provider.provider_fee + ) + build_error = None + except Exception as exc: + configured = None + build_error = f"{type(exc).__name__}: {exc}" + + served = _served_model_for_provider(row.id, provider_pk) + return _ModelEvaluation( + model_id=row.id, configured=configured, build_error=build_error, served=served + ) + + +def _report_row( + row_id: str, status: str, title: str, detail: str, evidence: dict[str, object] +) -> dict[str, object]: + return { + "id": row_id, + "status": status, + "title": title, + "detail": detail, + "evidence": evidence, + } + + +def _aggregate_row( + row_id: str, + title: str, + checked: int, + flagged: Sequence[object], + *, + fail_status: str, + empty_detail: str, + ok_detail: str, + flagged_detail: str, +) -> dict[str, object]: + """The ok/fail(-or-warn) shape every pricing row shares: examine + ``checked`` items, flag some of them as a problem, report the count. + """ + evidence: dict[str, object] = {"checked": checked, "flagged": list(flagged)} + if checked == 0: + return _report_row(row_id, "ok", title, empty_detail, evidence) + if flagged: + return _report_row(row_id, fail_status, title, flagged_detail, evidence) + return _report_row(row_id, "ok", title, ok_detail, evidence) + + +def _report_row_served_matches_configured( + evaluations: list[_ModelEvaluation], +) -> dict[str, object]: + mismatched: list[dict[str, object]] = [] + for ev in evaluations: + if ev.configured is None: + mismatched.append( + { + "model_id": ev.model_id, + "configured": None, + "served": None, + "error": ev.build_error, + } + ) + continue + + served_pricing = ev.served.pricing.dict() if ev.served else None + if served_pricing != ev.configured.pricing.dict(): + mismatched.append( + { + "model_id": ev.model_id, + "configured": ev.configured.pricing.dict(), + "served": served_pricing, + } + ) + + checked = len(evaluations) + return _aggregate_row( + "pricing.served_matches_configured", + "Served price matches configured price", + checked, + mismatched, + fail_status="fail", + empty_detail="No enabled models to check.", + ok_detail=f"All {checked} enabled models are served at the configured price.", + flagged_detail=( + f"{len(mismatched)} of {checked} enabled models have a served price " + "that disagrees with the configured price." + ), + ) + + +def _report_row_sats_pricing_present( + evaluations: list[_ModelEvaluation], +) -> dict[str, object]: + checked = 0 + missing: list[str] = [] + for ev in evaluations: + if ev.served is None: + continue + checked += 1 + if ev.served.sats_pricing is None: + missing.append(ev.model_id) + + return _aggregate_row( + "pricing.sats_pricing_present", + "Sats pricing computed for served models", + checked, + missing, + fail_status="fail", + empty_detail="No served models to check.", + ok_detail=f"All {checked} served models have a computed sats price.", + flagged_detail=f"{len(missing)} of {checked} served models have no computed sats price.", + ) + + +def _report_row_enabled_models_served( + evaluations: list[_ModelEvaluation], +) -> dict[str, object]: + missing = [ev.model_id for ev in evaluations if ev.served is None] + checked = len(evaluations) + return _aggregate_row( + "pricing.enabled_models_served", + "Enabled models are served", + checked, + missing, + fail_status="fail", + empty_detail="No enabled models to check.", + ok_detail=f"All {checked} enabled models are being served.", + flagged_detail=f"{len(missing)} of {checked} enabled models are not being served.", + ) + + +def _report_row_cache_rate( + evaluations: list[_ModelEvaluation], +) -> dict[str, object]: + checked = 0 + unknown: list[dict[str, object]] = [] + for ev in evaluations: + # Scoped to served models only, like every sibling row: a model the + # routing algorithm withholds from the served map (e.g. an unusable + # stored price) has no cache-billing behaviour to certify here — + # ``pricing.enabled_models_served`` already flags it as unserved. + if ev.served is None or ev.configured is None: + continue + checked += 1 + + pricing = ev.configured.pricing + missing_rates = [] + if (pricing.input_cache_read or 0.0) <= 0.0: + missing_rates.append("input_cache_read") + if (pricing.input_cache_write or 0.0) <= 0.0: + missing_rates.append("input_cache_write") + if missing_rates: + unknown.append({"model_id": ev.model_id, "missing_rates": missing_rates}) + + return _aggregate_row( + "pricing.cache_rate", + "Cache rate known for served models", + checked, + unknown, + fail_status="warn", + empty_detail="No served models to check.", + ok_detail=( + f"All {checked} served models have known cache-read and cache-write rates." + ), + flagged_detail=( + f"{len(unknown)} of {checked} served models are missing a cache-read or " + "cache-write rate." + ), + ) + + +@admin_router.get( + "/api/upstream-providers/{provider_id}/report", + dependencies=[Depends(require_admin_api)], +) +async def get_upstream_provider_report(provider_id: str) -> dict[str, object]: + """Certification report for one configured upstream provider. + + The four pricing rows (``pricing.served_matches_configured``, + ``pricing.sats_pricing_present``, ``pricing.enabled_models_served``, + ``pricing.cache_rate``) are computed from the DB row plus the in-process + served map; none of them make a network call, so the ``GET`` never spends + and never blocks on an upstream. Each enabled row is evaluated once and + the result shared across all four rows, rather than every row re-walking + the served map and re-parsing the stored pricing on its own. + """ + async with create_session() as session: + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) + result = await session.exec( + select(ModelRow).where( + ModelRow.upstream_provider_id == provider_pk, + ModelRow.enabled, + ) + ) + enabled_rows = list(result.all()) + + evaluations = [ + _evaluate_model_row(row, provider, provider_pk) for row in enabled_rows + ] + + rows = [ + _report_row_served_matches_configured(evaluations), + _report_row_sats_pricing_present(evaluations), + _report_row_enabled_models_served(evaluations), + _report_row_cache_rate(evaluations), + ] + + return { + "provider_id": provider.id, + "generated_at": datetime.now(timezone.utc).isoformat(), + "rows": rows, + } + + class CreateAccountRequest(BaseModel): provider_type: str From c5e4f4caee1edf7792c5d68f7888880b6c2dd029 Mon Sep 17 00:00:00 2001 From: 9qeklajc <211699015+9qeklajc@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:59:55 +0000 Subject: [PATCH 03/18] feat(certification): comprehensive upstream certification harness MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Extends PR #717's row contract with live network probes and a standalone CLI runner. The four operator-facing goals are: - Heartbeat — endpoint responds and is online - Usage data — tokens and requests captured - Cost data — prompt and completion cost calculated - Pricing in /v1/models — cost updates reflected in models list New endpoint: POST /admin/api/upstream-providers/{id}/certify - Probes the upstream's /models and sends a 1-token completion - Runs the node's own cost engine on the real response - Never enters the billing path (no reservation, no Cashu, no wallet) - Returns pricing.* rows + live rows + a checklist with ok/warn/fail ticks New module: routstr/upstream/certification.py - Pure row builders for every verdict (ok/warn/fail), testable without a socket - Independent cost re-derivation: reproduces _calculate_from_tokens arithmetic rather than calling the engine and comparing it to itself - CLI runner: python -m routstr.upstream.certification --url Tests: 49 unit + 14 integration, all green. Co-Authored-By: PR #717 (jeroenubbink) for the row contract and pricing rows. --- routstr/core/admin.py | 141 +++ routstr/upstream/certification.py | 985 +++++++++++++++++++++ tests/integration/test_certify_endpoint.py | 600 +++++++++++++ tests/unit/test_certification.py | 605 +++++++++++++ 4 files changed, 2331 insertions(+) create mode 100644 routstr/upstream/certification.py create mode 100644 tests/integration/test_certify_endpoint.py create mode 100644 tests/unit/test_certification.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 290480ac..8d95d936 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1500,6 +1500,147 @@ async def get_upstream_provider_report(provider_id: str) -> dict[str, object]: } +class CertifyRequest(BaseModel): + model_id: str | None = None + timeout_seconds: float | None = None + + +@admin_router.post( + "/api/upstream-providers/{provider_id}/certify", + dependencies=[Depends(require_admin_api)], +) +async def certify_upstream_provider( + provider_id: str, payload: CertifyRequest +) -> dict[str, object]: + """Live certification checks for a configured upstream provider. + + Unlike the read-only ``GET …/report``, this endpoint probes the + upstream over the network: it calls ``/models`` and sends a one-token + completion, then runs the node's own cost engine on the real response. + It never enters the billing path — no reservation, no Cashu, no wallet + — so it cannot spend the node's wallet. It costs at most one + completion's worth of upstream credit. + + The response carries the four ``pricing.*`` rows from the read-only + report (re-derived here so the certification is self-contained) plus + the five live/derived rows from + :mod:`routstr.upstream.certification`, and a ``checklist`` summarising + the four operator-facing goals with ``ok``/``warn``/``fail`` ticks. + """ + from ..payment.price import sats_usd_price + from ..upstream.certification import ( + build_checklist, + run_live_checks, + ) + + async with create_session() as session: + provider = await _get_upstream_provider_by_ref(session, provider_id) + provider_pk = _provider_pk(provider) + result = await session.exec( + select(ModelRow).where( + ModelRow.upstream_provider_id == provider_pk, + ModelRow.enabled, + ) + ) + enabled_rows = list(result.all()) + + evaluations = [ + _evaluate_model_row(row, provider, provider_pk) for row in enabled_rows + ] + pricing_rows = [ + _report_row_served_matches_configured(evaluations), + _report_row_sats_pricing_present(evaluations), + _report_row_enabled_models_served(evaluations), + _report_row_cache_rate(evaluations), + ] + + model_id = payload.model_id + if not model_id and enabled_rows: + # Pick the first enabled row that is actually being served — a + # model withheld from the served map would fail the chat probe for + # a reason unrelated to the endpoint's health. + for ev in evaluations: + if ev.served is not None: + model_id = ev.served.id + break + if model_id is None: + model_id = enabled_rows[0].id + + from ..proxy import get_candidates + + model_obj = None + if model_id: + for model, _upstream in get_candidates(model_id) or []: + if model.upstream_provider_id == provider_pk: + model_obj = model + break + if model_obj is None: + from ..upstream.certification import ( + STATUS_WARN, + certification_row, + ) + + live_rows = [ + certification_row( + "endpoint.validity", + STATUS_WARN, + "Upstream URL is well-formed", + "No served model is available for this provider, so the " + "live checks could not run.", + {"base_url": provider.base_url}, + ), + certification_row( + "endpoint.reachable", + STATUS_WARN, + "Endpoint responds", + "Skipped — no model to probe.", + {}, + ), + certification_row( + "endpoint.models_payload", + STATUS_WARN, + "Models payload has the expected shape", + "Skipped — no model to probe.", + {}, + ), + certification_row( + "usage.capture", + STATUS_WARN, + "Token usage captured from a completion", + "Skipped — no model to probe.", + {}, + ), + certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + "Skipped — no model to probe.", + {}, + ), + ] + else: + sats_to_usd = sats_usd_price() + timeout = ( + payload.timeout_seconds if payload.timeout_seconds is not None else 15.0 + ) + live_rows = await run_live_checks( + provider.base_url, + provider.api_key, + model_obj, + provider_fee=provider.provider_fee, + sats_to_usd=sats_to_usd, + timeout=timeout, + ) + + rows = pricing_rows + live_rows + return { + "provider_id": provider.id, + "generated_at": datetime.now(timezone.utc).isoformat(), + "rows": rows, + "checklist": build_checklist(rows), + } + + class CreateAccountRequest(BaseModel): provider_type: str diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py new file mode 100644 index 00000000..6952a698 --- /dev/null +++ b/routstr/upstream/certification.py @@ -0,0 +1,985 @@ +"""Certification checks for an upstream provider endpoint. + +PR #717 established the row contract — ``{id, status, title, detail, +evidence}`` with ``status`` in ``{ok, warn, fail}`` — and the four pricing +rows derived from the database row plus the in-process served map. Those +rows deliberately never touch the network. This module adds the checks that +*must* touch the network, and the checklist view that maps the +operator-facing goals onto rows: + +========================= ========================================= +Goal Row(s) +========================= ========================================= +Heartbeat ``endpoint.reachable`` +Usage data ``usage.capture`` +Cost data ``cost.prompt_completion`` +Pricing in ``/v1/models`` ``pricing.served_matches_configured``, + ``pricing.enabled_models_served`` +========================= ========================================= + +**Money safety.** Every live check calls the upstream directly with +``httpx`` — exactly like the existing ``POST /api/models/test`` probe — and +never enters the node's billing path. No reservation is taken, no Cashu +token is minted or spent, and the probe asks for a single token +(``max_tokens=1``). A probe therefore costs the operator at most one +completion's worth of upstream spend and nothing from the node's wallet. + +**Why a separate endpoint.** ``GET …/report`` promises the operator a +cheap, non-blocking read. A live probe can hang for the length of its +timeout and spends upstream credit, so it lives behind +``POST …/certify`` instead of being folded into the read. +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import math +import time +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any +from urllib.parse import urlparse + +import httpx + +from ..core.logging import get_logger +from ..payment.cost_calculation import calculate_cost +from ..payment.usage import normalize_usage + +if TYPE_CHECKING: + from ..payment.models import Model + +logger = get_logger(__name__) + +STATUS_OK = "ok" +STATUS_WARN = "warn" +STATUS_FAIL = "fail" + +TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"} + +# A probe must never be able to wedge an admin request. Fifteen seconds is +# generous for a `/models` listing or a one-token completion on a healthy +# upstream, and bounded enough that a dead host fails the row rather than +# the request. +PROBE_TIMEOUT_SECONDS = 15.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. +PROBE_MAX_TOKENS = 1 +PROBE_PROMPT = "ping" + +# The reservation ceiling is irrelevant to the token-priced path — it is +# only the amount held before settlement — but ``calculate_cost`` requires +# one. Any value at or above the real charge behaves identically. +_PROBE_MAX_COST_MSATS = 1_000_000_000 + +# Rounding in ``_calculate_from_tokens`` truncates the output component and +# folds the remainder into the input component, so a one-millisatoshi +# difference is arithmetic, not drift. +COST_TOLERANCE_MSATS = 1 + + +def certification_row( + row_id: str, + status: str, + title: str, + detail: str, + evidence: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Build one row of the certification report.""" + return { + "id": row_id, + "status": status, + "title": title, + "detail": detail, + "evidence": evidence if evidence is not None else {}, + } + + +# 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 +# it ``warn``. Kept as data so the checklist and the row set cannot drift. +CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = ( + ( + "heartbeat", + "Heartbeat — endpoint responds and is online", + ("endpoint.reachable",), + ), + ( + "usage_data", + "Usage data — tokens and requests captured", + ("usage.capture",), + ), + ( + "cost_data", + "Cost data — prompt and completion cost calculated", + ("cost.prompt_completion",), + ), + ( + "pricing_v1_models", + "Pricing in /v1/models — cost updates reflected in the models list", + ("pricing.served_matches_configured", "pricing.enabled_models_served"), + ), +) + + +def build_checklist(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Summarise the rows as the four operator-facing goals with ticks.""" + by_id = {row["id"]: row for row in rows} + checklist: list[dict[str, Any]] = [] + for goal, label, row_ids in CHECKLIST_GOALS: + present = [by_id[row_id]["status"] for row_id in row_ids if row_id in by_id] + if not present: + status = STATUS_WARN + elif any(item == STATUS_FAIL for item in present): + status = STATUS_FAIL + elif all(item == STATUS_OK for item in present): + status = STATUS_OK + else: + status = STATUS_WARN + checklist.append( + { + "goal": goal, + "label": label, + "status": status, + "tick": TICKS[status], + "rows": [row_id for row_id in row_ids if row_id in by_id], + } + ) + return checklist + + +@dataclass +class ProbeResult: + """Raw outcome of the two live HTTP calls a probe makes.""" + + base_url: str + models_url: str + chat_url: str + models_status: int | None = None + models_payload: dict[str, Any] | None = None + models_error: str | None = None + models_latency_ms: float | None = None + chat_status: int | None = None + chat_payload: dict[str, Any] | None = None + chat_error: str | None = None + chat_latency_ms: float | None = None + + +async def probe_upstream( + base_url: str, + api_key: str, + model_id: str, + *, + client: httpx.AsyncClient | None = None, + timeout: float = PROBE_TIMEOUT_SECONDS, +) -> ProbeResult: + """Call the upstream's ``/models`` and a one-token completion. + + A transport failure on either call is recorded on the result rather + than raised: a dead upstream is a ``fail`` row, not a failed request. + """ + base = base_url.rstrip("/") + result = ProbeResult( + base_url=base_url, + models_url=f"{base}/models", + chat_url=f"{base}/chat/completions", + ) + headers = {"Content-Type": "application/json"} + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + + owns_client = client is None + if client is None: + client = httpx.AsyncClient(timeout=timeout) + + try: + started = time.monotonic() + try: + response = await client.get(result.models_url, headers=headers) + result.models_status = response.status_code + result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) + try: + body = response.json() + except Exception as exc: # noqa: BLE001 - any decode failure is the signal + result.models_error = f"{type(exc).__name__}: {exc}" + else: + if isinstance(body, dict): + result.models_payload = body + else: + result.models_error = ( + f"expected a JSON object, got {type(body).__name__}" + ) + except Exception as exc: # noqa: BLE001 - transport failure is a row status + result.models_error = f"{type(exc).__name__}: {exc}" + result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) + + started = time.monotonic() + if not model_id: + return result + request_body = { + "model": model_id, + "messages": [{"role": "user", "content": PROBE_PROMPT}], + "max_tokens": PROBE_MAX_TOKENS, + "stream": False, + } + try: + response = await client.post( + result.chat_url, json=request_body, headers=headers + ) + result.chat_status = response.status_code + result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) + try: + payload = response.json() + except Exception as exc: # noqa: BLE001 - any decode failure is the signal + result.chat_error = f"{type(exc).__name__}: {exc}" + else: + if isinstance(payload, dict): + result.chat_payload = payload + else: + result.chat_error = ( + f"expected a JSON object, got {type(payload).__name__}" + ) + except Exception as exc: # noqa: BLE001 - transport failure is a row status + result.chat_error = f"{type(exc).__name__}: {exc}" + result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) + finally: + if owns_client: + await client.aclose() + + return result + + +# --------------------------------------------------------------------------- +# Row builders +# +# Every builder below is pure: it turns an already-fetched fact (a probe +# result, a model, a computed cost) into a row. The network lives only in +# ``probe_upstream`` and ``run_live_checks``, so a test can exercise each +# verdict — including the failure ones — without a socket. +# --------------------------------------------------------------------------- + + +def endpoint_validity_row(base_url: str) -> dict[str, Any]: + """Check the configured base URL is a well-formed http(s) endpoint.""" + parsed = urlparse(base_url or "") + 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: + problems.append("no host component") + evidence: dict[str, Any] = { + "base_url": base_url, + "scheme": parsed.scheme, + "host": parsed.netloc, + "path": parsed.path, + } + if problems: + return certification_row( + "endpoint.validity", + STATUS_FAIL, + "Upstream URL is well-formed", + "The configured base URL is not a usable http(s) endpoint: " + + "; ".join(problems) + + ".", + evidence, + ) + return certification_row( + "endpoint.validity", + STATUS_OK, + "Upstream URL is well-formed", + f"{parsed.scheme}://{parsed.netloc} is a valid endpoint.", + evidence, + ) + + +def heartbeat_row(probe: ProbeResult) -> dict[str, Any]: + """Check the upstream's ``/models`` responds — the heartbeat.""" + evidence: dict[str, Any] = { + "url": probe.models_url, + "status_code": probe.models_status, + "latency_ms": probe.models_latency_ms, + } + if probe.models_status is None: + evidence["error"] = probe.models_error + return certification_row( + "endpoint.reachable", + STATUS_FAIL, + "Endpoint responds", + f"No response from {probe.models_url}: {probe.models_error}.", + evidence, + ) + if 200 <= probe.models_status < 300: + return certification_row( + "endpoint.reachable", + STATUS_OK, + "Endpoint responds", + f"{probe.models_url} answered {probe.models_status} in " + f"{probe.models_latency_ms} ms.", + evidence, + ) + return certification_row( + "endpoint.reachable", + STATUS_FAIL, + "Endpoint responds", + f"{probe.models_url} answered {probe.models_status}.", + evidence, + ) + + +def models_payload_row(probe: ProbeResult) -> dict[str, Any]: + """Check the ``/models`` payload matches the OpenAI list shape.""" + if probe.models_payload is None: + 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}.", + {"url": probe.models_url, "error": probe.models_error}, + ) + + data = probe.models_payload.get("data") + if not isinstance(data, list): + return certification_row( + "endpoint.models_payload", + STATUS_FAIL, + "Models payload has the expected shape", + f'Expected a top-level "data" list, got {type(data).__name__}.', + { + "url": probe.models_url, + "top_level_keys": sorted(probe.models_payload.keys()), + }, + ) + + ids = [ + item.get("id") + for item in data + if isinstance(item, dict) and isinstance(item.get("id"), str) + ] + evidence: dict[str, Any] = { + "url": probe.models_url, + "model_count": len(data), + "usable_ids": len(ids), + "sample_ids": ids[:5], + } + if not ids: + return certification_row( + "endpoint.models_payload", + STATUS_FAIL, + "Models payload has the expected shape", + f'The "data" list carries no entry with a string "id" ' + f"({len(data)} entries).", + evidence, + ) + return certification_row( + "endpoint.models_payload", + STATUS_OK, + "Models payload has the expected shape", + f"{len(ids)} of {len(data)} entries carry a string id.", + evidence, + ) + + +def usage_capture_row(probe: ProbeResult) -> dict[str, Any]: + """Check a completion comes back with token usage the node can bill on. + + A missing ``usage`` object is the root of the ``(0+0)`` billing bug — + the node has nothing to price, so the request settles for free. That is + a real defect in the upstream's OpenAI compatibility, but it does not + make the endpoint unusable, so it is a ``warn`` rather than a ``fail``. + """ + evidence: dict[str, Any] = { + "url": probe.chat_url, + "status_code": probe.chat_status, + "latency_ms": probe.chat_latency_ms, + } + if probe.chat_status is None: + evidence["error"] = probe.chat_error + return certification_row( + "usage.capture", + STATUS_FAIL, + "Token usage captured from a completion", + f"No response from {probe.chat_url}: {probe.chat_error}.", + evidence, + ) + if not 200 <= probe.chat_status < 300: + evidence["body"] = _truncate(probe.chat_payload) + return certification_row( + "usage.capture", + STATUS_FAIL, + "Token usage captured from a completion", + f"{probe.chat_url} answered {probe.chat_status} for a " + f"{PROBE_MAX_TOKENS}-token probe.", + evidence, + ) + if probe.chat_payload is None: + 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}.", + evidence, + ) + + raw_usage = probe.chat_payload.get("usage") + normalized = normalize_usage(raw_usage) + evidence["usage"] = raw_usage + if normalized is None: + return certification_row( + "usage.capture", + STATUS_WARN, + "Token usage captured from a completion", + 'The completion carried no "usage" object, so the node has no ' + "token counts to bill on and the request would settle as (0+0).", + evidence, + ) + evidence["input_tokens"] = normalized.input_tokens + evidence["output_tokens"] = normalized.output_tokens + if normalized.input_tokens <= 0 and normalized.output_tokens <= 0: + return certification_row( + "usage.capture", + STATUS_WARN, + "Token usage captured from a completion", + "The completion reported a usage object with zero tokens in both " + "directions.", + evidence, + ) + return certification_row( + "usage.capture", + STATUS_OK, + "Token usage captured from a completion", + f"Captured {normalized.input_tokens} input and " + f"{normalized.output_tokens} output tokens.", + evidence, + ) + + +def _truncate(value: Any, limit: int = 400) -> Any: + """Clip an upstream body so one bad response cannot bloat the report.""" + if value is None: + return None + text = value if isinstance(value, str) else json.dumps(value, default=str) + return text if len(text) <= limit else text[:limit] + "…" + + +def _reported_usd_cost(payload: dict[str, Any]) -> float: + """The upstream-reported USD cost, or 0.0 when it reported none. + + 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. + """ + 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) + 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) + return 0.0 + + +def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int]: + """Re-derive the token-priced charge independently of the engine. + + ``_calculate_from_tokens`` prices at *msats per 1000 tokens*, rounds + each component to three decimals, ceilings the sum, then folds the + cache cost into the input component by truncating the output one. The + arithmetic is reproduced here — rather than calling the engine and + comparing it to itself — so a swapped input/output rate, a dropped + cache term or a changed rounding rule shows up as a mismatch. + + Returns ``(total_msats, input_msats, output_msats)``. + """ + input_rate = float(sats_pricing.prompt) * 1_000_000.0 + output_rate = float(sats_pricing.completion) * 1_000_000.0 + cache_read_rate = ( + float(sats_pricing.input_cache_read or 0.0) * 1_000_000.0 or input_rate + ) + cache_write_rate = ( + float(sats_pricing.input_cache_write or 0.0) * 1_000_000.0 or input_rate + ) + + 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) + calc_cache_write = round(usage.cache_write_tokens / 1000 * cache_write_rate, 3) + + total = math.ceil(calc_input + calc_output + calc_cache_read + calc_cache_write) + visible_output = int(calc_output) + return total, total - visible_output, visible_output + + +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.""" + return math.ceil(reported_usd * provider_fee / sats_to_usd * 1000) + + +def cost_prompt_completion_row( + *, + model: "Model", + probe: ProbeResult, + cost_data: Any, + provider_fee: float, + sats_to_usd: float, + pricing_known: bool = True, +) -> dict[str, Any]: + """Check the node's cost engine prices a real completion correctly. + + Both the prompt and the completion component are checked: the engine + truncates the output component and folds the remainder into the input + component so that ``input + output == total`` exactly, which means a + wrong rate on *either* side shows up as a mismatch here. + """ + from ..payment.cost_calculation import CostDataError + + payload = probe.chat_payload or {} + usage = normalize_usage(payload.get("usage")) + evidence: dict[str, Any] = { + "model_id": model.id, + "forwarded_model_id": model.forwarded_model_id, + "provider_fee": provider_fee, + "sats_usd_price": sats_to_usd, + } + + if isinstance(cost_data, CostDataError): + evidence["error"] = cost_data.message + return certification_row( + "cost.prompt_completion", + STATUS_FAIL, + "Prompt and completion cost calculated", + f"The cost engine could not price the completion: {cost_data.message}.", + evidence, + ) + if usage is None: + return certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + "No token usage to price — see the usage row.", + evidence, + ) + if model.sats_pricing is None: + return certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + "This model has no computed sats pricing, so there is nothing to " + "verify the charge against.", + evidence, + ) + if not pricing_known: + return certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + "No pricing is known for this model, so the charge cannot be " + "verified. Configure the model on the node, or pass explicit " + "prices, to certify this row.", + evidence, + ) + + 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 + ) + basis = "configured_token_pricing" + + actual_total = int(cost_data.total_msats) + actual_input = int(cost_data.input_msats) + actual_output = int(cost_data.output_msats) + + evidence.update( + { + "basis": basis, + "reported_usd": reported_usd or None, + "input_tokens": usage.input_tokens, + "output_tokens": usage.output_tokens, + "cache_read_tokens": usage.cache_read_tokens, + "cache_write_tokens": usage.cache_write_tokens, + "expected_total_msats": expected_total, + "expected_input_msats": expected_input, + "expected_output_msats": expected_output, + "actual_total_msats": actual_total, + "actual_input_msats": actual_input, + "actual_output_msats": actual_output, + } + ) + + mismatches: list[str] = [] + if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS: + mismatches.append(f"total {actual_total} != {expected_total}") + if actual_input + actual_output != actual_total: + mismatches.append( + f"components {actual_input}+{actual_output} != total {actual_total}" + ) + if ( + expected_output is not None + and abs(actual_output - expected_output) > COST_TOLERANCE_MSATS + ): + mismatches.append(f"output {actual_output} != {expected_output}") + if ( + expected_input is not None + and abs(actual_input - expected_input) > COST_TOLERANCE_MSATS + ): + mismatches.append(f"input {actual_input} != {expected_input}") + + if mismatches: + return certification_row( + "cost.prompt_completion", + STATUS_FAIL, + "Prompt and completion cost calculated", + "The computed charge disagrees with the configured pricing: " + + "; ".join(mismatches) + + ".", + evidence, + ) + return certification_row( + "cost.prompt_completion", + STATUS_OK, + "Prompt and completion cost calculated", + f"Charged {actual_total} msats ({actual_input} input + " + f"{actual_output} output) for {usage.input_tokens} prompt and " + f"{usage.output_tokens} completion tokens, matching the configured " + f"pricing.", + evidence, + ) + + +async def run_live_checks( + base_url: str, + api_key: str, + model: "Model", + *, + provider_fee: float, + sats_to_usd: float, + client: httpx.AsyncClient | None = None, + timeout: float = PROBE_TIMEOUT_SECONDS, + pricing_known: bool = True, +) -> list[dict[str, Any]]: + """Probe one upstream once and build the five live/derived rows.""" + probe = await probe_upstream( + base_url, + api_key, + model.forwarded_model_id or model.id, + client=client, + timeout=timeout, + ) + rows = [ + endpoint_validity_row(base_url), + heartbeat_row(probe), + models_payload_row(probe), + usage_capture_row(probe), + ] + + cost_data: Any = None + if probe.chat_payload is not None and probe.chat_status is not None: + try: + cost_data = await calculate_cost( + probe.chat_payload, + _PROBE_MAX_COST_MSATS, + model_obj=model, + provider_fee=provider_fee, + ) + except Exception as exc: # noqa: BLE001 - a raising engine is a fail row + from ..payment.cost_calculation import CostDataError + + cost_data = CostDataError( + message=f"{type(exc).__name__}: {exc}", code="pricing_error" + ) + if cost_data is None: + from ..payment.cost_calculation import CostDataError + + cost_data = CostDataError( + message=probe.chat_error or "the completion probe did not succeed", + code="no_completion", + ) + + 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, + ) + ) + return rows + + +# --------------------------------------------------------------------------- +# Standalone runner +# +# ``certify_upstream_url`` deliberately reads nothing from the node's +# database: the point of the CLI is to certify a URL *before* it is +# configured, or one the operator does not want to write into the node at +# all. The four pricing rows therefore do not apply here — they compare a +# stored row against the served map, neither of which exists for a bare +# URL — and the cost row falls back to litellm's cost map (or explicit +# prices) instead of a configured row. +# --------------------------------------------------------------------------- + + +def _first_model_id(probe: ProbeResult) -> str | None: + data = (probe.models_payload or {}).get("data") + if not isinstance(data, list): + return None + for item in data: + if isinstance(item, dict): + model_id = item.get("id") + if isinstance(model_id, str) and model_id: + return model_id + return None + + +def _as_price(value: Any) -> float | None: + if isinstance(value, (int, float)) and math.isfinite(value) and value >= 0: + return float(value) + return None + + +def _model_from_usd_pricing( + model_id: str, prompt_usd: float, completion_usd: float, sats_to_usd: float +) -> "Model": + """A throwaway ``Model`` carrying just enough to exercise the cost engine.""" + from ..payment.models import ( + Architecture, + Model, + Pricing, + _update_model_sats_pricing, + ) + + model = Model( + id=model_id, + name=model_id, + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=prompt_usd, completion=completion_usd), + sats_pricing=None, + per_request_limits=None, + top_provider=None, + enabled=True, + upstream_provider_id=None, + canonical_slug=None, + ) + return _update_model_sats_pricing(model, sats_to_usd) + + +async def certify_upstream_url( + base_url: str, + *, + api_key: str = "", + model_id: str | None = None, + prompt_price: float | None = None, + completion_price: float | None = None, + provider_fee: float = 1.0, + timeout: float = PROBE_TIMEOUT_SECONDS, + 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() + target: dict[str, Any] = {"base_url": base_url, "model_id": model_id} + + if not model_id: + discovery = await probe_upstream( + base_url, api_key, "", client=client, timeout=timeout + ) + model_id = _first_model_id(discovery) + target["model_id"] = model_id + if model_id is None: + rows = [ + endpoint_validity_row(base_url), + heartbeat_row(discovery), + models_payload_row(discovery), + certification_row( + "usage.capture", + STATUS_FAIL, + "Token usage captured from a completion", + "No model id is available to probe: pass --model, or the " + "upstream must list at least one id.", + {"url": discovery.chat_url}, + ), + certification_row( + "cost.prompt_completion", + STATUS_FAIL, + "Prompt and completion cost calculated", + "No model id is available to price.", + {}, + ), + ] + return { + "target": target, + "rows": rows, + "checklist": build_checklist(rows), + } + + entry = litellm_cost_entry(model_id) or {} + resolved_prompt = ( + prompt_price + if prompt_price is not None + else _as_price(entry.get("input_cost_per_token")) + ) + resolved_completion = ( + 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 + target["prompt_price_usd"] = resolved_prompt + target["completion_price_usd"] = resolved_completion + + model = _model_from_usd_pricing( + model_id, resolved_prompt or 0.0, resolved_completion or 0.0, sats_to_usd + ) + rows = await run_live_checks( + base_url, + api_key, + model, + provider_fee=provider_fee, + sats_to_usd=sats_to_usd, + client=client, + timeout=timeout, + pricing_known=pricing_known, + ) + return {"target": target, "rows": rows, "checklist": build_checklist(rows)} + + +def render_checklist(result: dict[str, Any]) -> str: + """Render one certification result as the operator-facing checklist.""" + target = result.get("target", {}) + lines = [f"Upstream certification — {target.get('base_url')}"] + if target.get("model_id"): + lines.append(f" model: {target['model_id']}") + lines.append("") + lines.append(" checklist") + for item in result.get("checklist", []): + lines.append(f" {item['tick']} {item['label']}") + lines.append("") + lines.append(" rows") + for row in result.get("rows", []): + tick = TICKS.get(row["status"], "?") + lines.append(f" {tick} [{row['id']}] {row['detail']}") + return "\n".join(lines) + + +def main(argv: list[str] | None = None) -> int: + """Run the checklist against one or more upstream base URLs.""" + parser = argparse.ArgumentParser( + prog="python -m routstr.upstream.certification", + description=( + "Run the upstream certification checklist against one or more " + "upstream base URLs. Exits non-zero when any row fails." + ), + ) + parser.add_argument( + "--url", + action="append", + required=True, + help="Upstream base URL (repeatable), e.g. https://api.example.com/v1", + ) + parser.add_argument("--key", default="", help="Bearer API key for the upstream") + parser.add_argument( + "--model", + default=None, + help="Model id to probe (defaults to the first id the upstream lists)", + ) + parser.add_argument( + "--prompt-price", + type=float, + default=None, + help="USD per prompt token (defaults to litellm's cost map)", + ) + parser.add_argument( + "--completion-price", + type=float, + default=None, + help="USD per completion token (defaults to litellm's cost map)", + ) + parser.add_argument( + "--provider-fee", + type=float, + default=1.0, + help="Provider fee multiplier applied by the cost check", + ) + parser.add_argument( + "--timeout", + type=float, + default=PROBE_TIMEOUT_SECONDS, + help="Per-request probe timeout in seconds", + ) + parser.add_argument( + "--json", action="store_true", help="Emit the raw report as JSON" + ) + args = parser.parse_args(argv) + + async def _run_all() -> list[dict[str, Any]]: + results: list[dict[str, Any]] = [] + for url in args.url: + results.append( + await certify_upstream_url( + url, + api_key=args.key, + model_id=args.model, + prompt_price=args.prompt_price, + completion_price=args.completion_price, + provider_fee=args.provider_fee, + timeout=args.timeout, + ) + ) + return results + + results = asyncio.run(_run_all()) + + if args.json: + print(json.dumps(results, indent=2, default=str)) + else: + for result in results: + print(render_checklist(result)) + print() + + worst = STATUS_OK + for result in results: + for row in result.get("rows", []): + if row["status"] == STATUS_FAIL: + worst = STATUS_FAIL + elif row["status"] == STATUS_WARN and worst == STATUS_OK: + worst = STATUS_WARN + return 1 if worst == STATUS_FAIL else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py new file mode 100644 index 00000000..3f5b040f --- /dev/null +++ b/tests/integration/test_certify_endpoint.py @@ -0,0 +1,600 @@ +"""Integration tests for POST /admin/api/upstream-providers/{id}/certify. + +Exercises the endpoint against a mocked upstream (no real network, no +real spend). The test fixtures create a provider + model row in the +integration DB, then mock the two HTTP calls the probe makes (GET /models +and POST /chat/completions) with ``respx`` so every verdict — ok, warn, +fail — is reachable deterministically. +""" + +from __future__ import annotations + +import json +from datetime import datetime, timedelta, timezone +from typing import Any +from unittest.mock import patch + +import pytest +import respx +from httpx import AsyncClient, Response +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + + +# The conftest patches ``routstr.payment.price.sats_usd_price``, but +# ``cost_calculation.py`` and ``models.py`` import it as a local binding +# the conftest-level patch cannot reach. Pin it here so every test that +# goes through ``_row_to_model`` or ``calculate_cost`` gets a real sats +# price — same pattern as ``test_model_price_propagation.py``. +@pytest.fixture(autouse=True) +def _pin_sats_usd() -> Any: + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + with patch( + "routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005 + ): + with patch("routstr.payment.price.SATS_USD_PRICE", 0.0005): + yield + + +ARCHITECTURE = { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, +} + + +def _pricing(**overrides: float) -> dict[str, Any]: + pricing: dict[str, Any] = { + "prompt": 1.4e-7, + "completion": 2.8e-7, + "request": 0.0, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, + "input_cache_read": 0.0, + "input_cache_write": 0.0, + } + pricing.update(overrides) + return pricing + + +def _admin_headers() -> dict[str, str]: + token = "test-certify-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +async def _make_provider( + session: AsyncSession, + *, + slug: str | None = None, + provider_fee: float = 1.0, + base_url: str = "https://certify-upstream.example/v1", + api_key: str = "test-key", +) -> UpstreamProviderRow: + provider = UpstreamProviderRow( + provider_type="generic", + base_url=base_url, + api_key=api_key, + provider_fee=provider_fee, + slug=slug, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + assert provider.id is not None + return provider + + +def _model_row(provider_id: int, **overrides: Any) -> ModelRow: + model_id = overrides.pop("model_id", "cert-test-model") + return ModelRow( + id=model_id, + name=model_id, + description="d", + created=0, + context_length=8192, + architecture=json.dumps(ARCHITECTURE), + pricing=json.dumps(_pricing(**overrides.pop("pricing_overrides", {}))), + upstream_provider_id=provider_id, + enabled=True, + forwarded_model_id=model_id, + ) + + +async def _seed_and_init( + session: AsyncSession, + client: AsyncClient, + *, + provider_fee: float = 1.0, + model_id: str = "cert-test-model", + pricing_overrides: dict[str, Any] | None = None, + base_url: str = "https://certify-upstream.example/v1", +) -> int: + provider = await _make_provider( + session, provider_fee=provider_fee, base_url=base_url + ) + session.add( + _model_row( + provider.id, + model_id=model_id, + pricing_overrides=pricing_overrides or {}, + ) + ) + await session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + return provider.id + + +def _find_row(rows: list[dict[str, Any]], row_id: str) -> dict[str, Any]: + for row in rows: + if row["id"] == row_id: + return row + raise AssertionError(f"row {row_id!r} not found") + + +def _mock_models_response( + base_url: str = "https://certify-upstream.example/v1", + models: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + if models is None: + models = [{"id": "cert-test-model"}] + return {"data": models} + + +def _mock_chat_response( + model: str = "cert-test-model", + prompt_tokens: int = 5, + completion_tokens: int = 1, +) -> dict[str, Any]: + return { + "id": "chatcmpl-test", + "object": "chat.completion", + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + }, + } + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_certify_requires_admin_auth( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider = await _make_provider(integration_session) + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", json={} + ) + assert resp.status_code == 403 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_certify_unknown_provider_returns_404( + integration_client: AsyncClient, +) -> None: + resp = await integration_client.post( + "/admin/api/upstream-providers/999999999/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 404 + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_all_ok( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + assert "rows" in body + assert "checklist" in body + + # All live rows should be ok + live_row_ids = [ + "endpoint.validity", + "endpoint.reachable", + "endpoint.models_payload", + "usage.capture", + "cost.prompt_completion", + ] + for row_id in live_row_ids: + row = _find_row(body["rows"], row_id) + assert row["status"] == "ok", f"{row_id}: {row}" + + # All checklist goals should be ok + for item in body["checklist"]: + assert item["status"] == "ok", f"{item['goal']}: {item}" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_heartbeat_fail_on_500( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(500, json={"error": "internal"}) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row = _find_row(body["rows"], "endpoint.reachable") + assert row["status"] == "fail" + assert row["evidence"]["status_code"] == 500 + + # heartbeat goal should be fail + heartbeat_goal = next( + item for item in body["checklist"] if item["goal"] == "heartbeat" + ) + assert heartbeat_goal["status"] == "fail" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_heartbeat_fail_on_transport_error( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + side_effect=__import__("httpx").ConnectError("connection refused") + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row = _find_row(body["rows"], "endpoint.reachable") + assert row["status"] == "fail" + assert row["evidence"]["error"] is not None + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_models_payload_fail_on_malformed( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json={"error": "no data field"}) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row = _find_row(body["rows"], "endpoint.models_payload") + assert row["status"] == "fail" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_usage_warn_when_no_usage( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + # No "usage" key in the chat response + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response( + 200, + json={ + "id": "x", + "model": "cert-test-model", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + }, + ) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + usage_row = _find_row(body["rows"], "usage.capture") + assert usage_row["status"] == "warn" + + cost_row = _find_row(body["rows"], "cost.prompt_completion") + assert cost_row["status"] == "warn" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_usage_fail_on_non_2xx( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(401, json={"error": {"message": "invalid api key"}}) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row = _find_row(body["rows"], "usage.capture") + assert row["status"] == "fail" + assert "401" in row["detail"] + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_cost_ok_with_token_pricing( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init( + integration_session, + integration_client, + provider_fee=1.0, + pricing_overrides={"prompt": 1e-7, "completion": 2e-7}, + ) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response( + 200, json=_mock_chat_response(prompt_tokens=10, completion_tokens=5) + ) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row = _find_row(body["rows"], "cost.prompt_completion") + assert row["status"] == "ok", row + assert row["evidence"]["basis"] == "configured_token_pricing" + assert row["evidence"]["input_tokens"] == 10 + assert row["evidence"]["output_tokens"] == 5 + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_cost_ok_with_usd_reported( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init( + integration_session, + integration_client, + provider_fee=1.05, + ) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + chat_payload = _mock_chat_response() + chat_payload["usage"]["cost_details"] = {"total_cost": 0.0001} + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=chat_payload) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row = _find_row(body["rows"], "cost.prompt_completion") + assert row["status"] == "ok", row + assert row["evidence"]["basis"] == "upstream_reported_usd" + assert row["evidence"]["reported_usd"] == 0.0001 + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_certify_with_no_served_model( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """When the provider has no model that is being served (e.g. all have + unusable pricing), the live checks should be skipped as warn, not + crash.""" + provider = await _make_provider(integration_session) + # A negative prompt price makes has_usable_pricing() return False, + # withholding the model from the served map. + session_add = _model_row( + provider.id, model_id="bad-model", pricing_overrides={"prompt": -1.0} + ) + integration_session.add(session_add) + await integration_session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + # Live rows should be warn (skipped) + for row_id in ["endpoint.reachable", "usage.capture", "cost.prompt_completion"]: + row = _find_row(body["rows"], row_id) + assert row["status"] == "warn", f"{row_id}: {row}" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_with_explicit_model_id( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init( + integration_session, + integration_client, + model_id="explicit-model", + ) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response( + 200, json=_mock_models_response(models=[{"id": "explicit-model"}]) + ) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response(model="explicit-model")) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={"model_id": "explicit-model"}, + ) + assert resp.status_code == 200, resp.text + body = resp.json() + usage_row = _find_row(body["rows"], "usage.capture") + assert usage_row["status"] == "ok" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_includes_pricing_rows( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """The certify response should also carry the four pricing.* rows from + the read-only report, so the certification is self-contained.""" + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + row_ids = [row["id"] for row in body["rows"]] + assert "pricing.served_matches_configured" in row_ids + assert "pricing.sats_pricing_present" in row_ids + assert "pricing.enabled_models_served" in row_ids + assert "pricing.cache_rate" in row_ids + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_row_contract_shape( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """Every row in the response has the required keys and a valid status.""" + provider_id = await _seed_and_init(integration_session, integration_client) + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200 + body = resp.json() + + assert "provider_id" in body + assert "generated_at" in body + assert "rows" in body + assert "checklist" in body + + for row in body["rows"]: + assert set(row) >= {"id", "status", "title", "detail", "evidence"} + assert row["status"] in {"ok", "warn", "fail"} + + for item in body["checklist"]: + assert set(item) >= {"goal", "label", "status", "tick", "rows"} + assert item["status"] in {"ok", "warn", "fail"} + assert item["tick"] in {"☑️", "⚠️", "❌"} diff --git a/tests/unit/test_certification.py b/tests/unit/test_certification.py new file mode 100644 index 00000000..dc3ae9fc --- /dev/null +++ b/tests/unit/test_certification.py @@ -0,0 +1,605 @@ +"""Unit tests for the pure row builders in routstr.upstream.certification. + +Every builder here takes an already-fetched fact (a ``ProbeResult``, a +model, a cost datum) and turns it into a row — no network, no DB. The +tests therefore cover every verdict — ok, warn, fail — including the +failure modes that would be hard to provoke against a live upstream. +""" + +from __future__ import annotations + +from typing import Any + +from routstr.upstream.certification import ( + STATUS_FAIL, + STATUS_OK, + STATUS_WARN, + ProbeResult, + _expected_token_msats, + _reported_usd_cost, + build_checklist, + certification_row, + cost_prompt_completion_row, + endpoint_validity_row, + heartbeat_row, + models_payload_row, + usage_capture_row, +) + + +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"), + models_latency_ms=kwargs.get("models_latency_ms", 42.0), + chat_status=kwargs.get("chat_status"), + chat_payload=kwargs.get("chat_payload"), + chat_error=kwargs.get("chat_error"), + chat_latency_ms=kwargs.get("chat_latency_ms", 88.0), + ) + + +# --- endpoint_validity_row -------------------------------------------------- + + +class TestEndpointValidity: + def test_ok_for_https_url(self) -> None: + row = endpoint_validity_row("https://api.example.com/v1") + assert row["status"] == STATUS_OK + assert row["evidence"]["scheme"] == "https" + assert row["evidence"]["host"] == "api.example.com" + + def test_ok_for_http_url(self) -> None: + row = endpoint_validity_row("http://localhost:8888/v1") + assert row["status"] == STATUS_OK + + def test_fail_for_ftp_scheme(self) -> None: + row = endpoint_validity_row("ftp://files.example.com") + assert row["status"] == STATUS_FAIL + assert "scheme" in row["detail"] + + def test_fail_for_empty_string(self) -> None: + row = endpoint_validity_row("") + assert row["status"] == STATUS_FAIL + + def test_fail_for_no_host(self) -> None: + row = endpoint_validity_row("https://") + assert row["status"] == STATUS_FAIL + + +# --- heartbeat_row ---------------------------------------------------------- + + +class TestHeartbeat: + def test_ok_on_2xx(self) -> None: + row = heartbeat_row(_probe(models_status=200)) + assert row["status"] == STATUS_OK + assert "200" in row["detail"] + assert row["evidence"]["latency_ms"] == 42.0 + + def test_ok_on_201(self) -> None: + row = heartbeat_row(_probe(models_status=201)) + assert row["status"] == STATUS_OK + + def test_fail_on_404(self) -> None: + row = heartbeat_row(_probe(models_status=404)) + assert row["status"] == STATUS_FAIL + assert "404" in row["detail"] + + def test_fail_on_500(self) -> None: + row = heartbeat_row(_probe(models_status=500)) + assert row["status"] == STATUS_FAIL + + def test_fail_on_transport_error(self) -> None: + row = heartbeat_row( + _probe(models_status=None, models_error="ConnectError: ...") + ) + assert row["status"] == STATUS_FAIL + assert "ConnectError" in row["detail"] + + +# --- models_payload_row ----------------------------------------------------- + + +class TestModelsPayload: + def test_ok_with_ids(self) -> None: + row = models_payload_row( + _probe(models_payload={"data": [{"id": "gpt-4o"}, {"id": "claude"}]}) + ) + assert row["status"] == STATUS_OK + assert row["evidence"]["model_count"] == 2 + assert row["evidence"]["usable_ids"] == 2 + assert row["evidence"]["sample_ids"] == ["gpt-4o", "claude"] + + def test_fail_when_data_missing(self) -> None: + row = models_payload_row(_probe(models_payload={"error": "not found"})) + assert row["status"] == STATUS_FAIL + assert "data" in row["detail"] + + def test_fail_when_data_not_list(self) -> None: + row = models_payload_row(_probe(models_payload={"data": {"id": "oops"}})) + assert row["status"] == STATUS_FAIL + + def test_fail_when_no_string_ids(self) -> None: + row = models_payload_row( + _probe(models_payload={"data": [{"name": "no-id-here"}]}) + ) + assert row["status"] == STATUS_FAIL + assert row["evidence"]["model_count"] == 1 + assert row["evidence"]["usable_ids"] == 0 + + def test_fail_when_payload_none(self) -> None: + row = models_payload_row( + _probe(models_payload=None, models_error="JSONDecodeError") + ) + assert row["status"] == STATUS_FAIL + + def test_sample_ids_truncated_to_five(self) -> None: + row = models_payload_row( + _probe(models_payload={"data": [{"id": f"m{i}"} for i in range(10)]}) + ) + assert len(row["evidence"]["sample_ids"]) == 5 + assert row["evidence"]["model_count"] == 10 + + +# --- usage_capture_row ------------------------------------------------------ + + +class TestUsageCapture: + def test_ok_with_tokens(self) -> None: + row = usage_capture_row( + _probe( + chat_status=200, + chat_payload={ + "choices": [], + "usage": {"prompt_tokens": 5, "completion_tokens": 1}, + }, + ) + ) + assert row["status"] == STATUS_OK + assert row["evidence"]["input_tokens"] == 5 + assert row["evidence"]["output_tokens"] == 1 + + def test_warn_when_no_usage_object(self) -> None: + row = usage_capture_row(_probe(chat_status=200, chat_payload={"choices": []})) + assert row["status"] == STATUS_WARN + assert "usage" in row["detail"].lower() + + def test_warn_when_all_zero_tokens(self) -> None: + row = usage_capture_row( + _probe( + chat_status=200, + chat_payload={ + "usage": {"prompt_tokens": 0, "completion_tokens": 0}, + }, + ) + ) + assert row["status"] == STATUS_WARN + + def test_fail_on_non_2xx(self) -> None: + row = usage_capture_row( + _probe(chat_status=401, chat_payload={"error": "bad key"}) + ) + assert row["status"] == STATUS_FAIL + assert "401" in row["detail"] + + def test_fail_on_transport_error(self) -> None: + row = usage_capture_row(_probe(chat_status=None, chat_error="TimeoutException")) + assert row["status"] == STATUS_FAIL + + def test_fail_when_body_not_json(self) -> None: + row = usage_capture_row( + _probe(chat_status=200, chat_payload=None, chat_error="not json") + ) + assert row["status"] == STATUS_FAIL + + def test_ok_with_anthropic_style_usage(self) -> None: + """Anthropic reports input_tokens (not prompt_tokens) and caches + additively, not as a grand total. normalize_usage handles this.""" + row = usage_capture_row( + _probe( + chat_status=200, + chat_payload={ + "usage": { + "input_tokens": 10, + "output_tokens": 2, + "cache_read_input_tokens": 5, + "cache_creation_input_tokens": 3, + }, + }, + ) + ) + assert row["status"] == STATUS_OK + + +# --- _reported_usd_cost ----------------------------------------------------- + + +class TestReportedUsdCost: + def test_zero_when_no_usage(self) -> None: + assert _reported_usd_cost({}) == 0.0 + assert _reported_usd_cost({"usage": None}) == 0.0 + + def test_from_cost_details_total(self) -> None: + payload = {"usage": {"cost_details": {"total_cost": 0.001}}} + assert _reported_usd_cost(payload) == 0.001 + + def test_from_total_cost(self) -> None: + payload = {"usage": {"total_cost": 0.002}} + assert _reported_usd_cost(payload) == 0.002 + + def test_from_cost_field(self) -> None: + payload = {"usage": {"cost": 0.003}} + assert _reported_usd_cost(payload) == 0.003 + + def test_zero_for_negative(self) -> None: + payload = {"usage": {"cost": -1.0}} + assert _reported_usd_cost(payload) == 0.0 + + def test_zero_for_nan(self) -> None: + payload = {"usage": {"cost": float("nan")}} + assert _reported_usd_cost(payload) == 0.0 + + def test_zero_for_non_numeric(self) -> None: + payload = {"usage": {"cost": "free"}} + assert _reported_usd_cost(payload) == 0.0 + + +# --- _expected_token_msats -------------------------------------------------- + + +class TestExpectedTokenMsats: + def _pricing(self, **overrides: float) -> Any: + from routstr.payment.models import Pricing + + pricing = Pricing(prompt=1.4e-7, completion=2.8e-7) + return pricing.copy(update=overrides) + + def _usage(self, **kwargs: int) -> Any: + from routstr.payment.usage import NormalizedUsage + + return NormalizedUsage(**kwargs) + + def test_basic_calculation(self) -> None: + pricing = self._pricing() + usage = self._usage(input_tokens=100, output_tokens=50) + total, inp, outp = _expected_token_msats(pricing, usage) + assert total > 0 + assert inp + outp == total # folding invariant + + def test_zero_tokens_give_zero_total(self) -> None: + pricing = self._pricing() + usage = self._usage() + total, inp, outp = _expected_token_msats(pricing, usage) + assert total == 0 + assert inp == 0 + assert outp == 0 + + def test_cache_tokens_included_in_total(self) -> None: + pricing = self._pricing(input_cache_read=0.5e-7, input_cache_write=0.7e-7) + usage = self._usage( + input_tokens=10, output_tokens=5, cache_read_tokens=3, cache_write_tokens=2 + ) + total, _, _ = _expected_token_msats(pricing, usage) + assert total > 0 + + def test_input_plus_output_equals_total(self) -> None: + """The folding invariant: visible_input = total - visible_output.""" + pricing = self._pricing(prompt=3.33e-7, completion=7.77e-7) + usage = self._usage(input_tokens=77, output_tokens=33) + total, inp, outp = _expected_token_msats(pricing, usage) + assert inp + outp == total + + +# --- cost_prompt_completion_row --------------------------------------------- + + +class TestCostPromptCompletion: + def _model(self, prompt: float = 1.4e-7, completion: float = 2.8e-7) -> Any: + from routstr.payment.models import ( + Architecture, + Model, + Pricing, + _update_model_sats_pricing, + ) + + model = Model( + id="test-model", + name="test-model", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=prompt, completion=completion), + sats_pricing=None, + per_request_limits=None, + top_provider=None, + enabled=True, + upstream_provider_id=None, + canonical_slug=None, + ) + return _update_model_sats_pricing(model, 0.0005) + + def _cost_data(self, total: int, inp: int, outp: int) -> Any: + from routstr.payment.cost_calculation import CostData + + return CostData( + base_msats=0, + input_msats=inp, + output_msats=outp, + total_msats=total, + ) + + def test_ok_when_engine_matches_expected(self) -> None: + model = self._model() + usage_dict = {"prompt_tokens": 10, "completion_tokens": 5} + probe = _probe(chat_status=200, chat_payload={"usage": usage_dict}) + + from routstr.payment.usage import normalize_usage + + usage = normalize_usage(usage_dict) + expected_total, expected_input, expected_output = _expected_token_msats( + model.sats_pricing, usage + ) + + cost_data = self._cost_data(expected_total, expected_input, expected_output) + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_OK + + def test_fail_when_total_mismatches(self) -> None: + model = self._model() + probe = _probe( + chat_status=200, + chat_payload={"usage": {"prompt_tokens": 10, "completion_tokens": 5}}, + ) + cost_data = self._cost_data(total=999, inp=500, outp=499) + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_FAIL + assert "total" in row["detail"] + + def test_fail_when_components_dont_sum(self) -> None: + model = self._model() + probe = _probe( + chat_status=200, + chat_payload={"usage": {"prompt_tokens": 10, "completion_tokens": 5}}, + ) + cost_data = self._cost_data(total=100, inp=60, outp=50) # 60+50 != 100 + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_FAIL + assert "components" in row["detail"] + + def test_warn_when_no_usage(self) -> None: + model = self._model() + probe = _probe(chat_status=200, chat_payload={"choices": []}) + cost_data = self._cost_data(total=0, inp=0, outp=0) + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_WARN + + def test_warn_when_no_sats_pricing(self) -> None: + from routstr.payment.models import Architecture, Model, Pricing + + model = Model( + id="no-sats", + name="no-sats", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=1e-7, completion=2e-7), + sats_pricing=None, + per_request_limits=None, + top_provider=None, + enabled=True, + upstream_provider_id=None, + canonical_slug=None, + ) + probe = _probe( + chat_status=200, + chat_payload={"usage": {"prompt_tokens": 5, "completion_tokens": 1}}, + ) + cost_data = self._cost_data(total=0, inp=0, outp=0) + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_WARN + + def test_warn_when_pricing_unknown(self) -> None: + model = self._model() + probe = _probe( + chat_status=200, + chat_payload={"usage": {"prompt_tokens": 5, "completion_tokens": 1}}, + ) + cost_data = self._cost_data(total=0, inp=0, outp=0) + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + pricing_known=False, + ) + assert row["status"] == STATUS_WARN + + def test_fail_on_cost_data_error(self) -> None: + from routstr.payment.cost_calculation import CostDataError + + model = self._model() + probe = _probe( + chat_status=200, + chat_payload={"usage": {"prompt_tokens": 5, "completion_tokens": 1}}, + ) + cost_data = CostDataError(message="pricing not found", code="pricing_error") + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_FAIL + assert "pricing not found" in row["detail"] + + def test_ok_with_usd_reported_cost(self) -> None: + model = self._model() + sats_to_usd = 0.0005 + provider_fee = 1.05 + reported_usd = 0.0001 + expected_total = int( + __import__("math").ceil(reported_usd * provider_fee / sats_to_usd * 1000) + ) + probe = _probe( + chat_status=200, + chat_payload={ + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "cost_details": {"total_cost": reported_usd}, + }, + }, + ) + cost_data = self._cost_data(expected_total, 0, expected_total) + row = cost_prompt_completion_row( + model=model, + probe=probe, + cost_data=cost_data, + provider_fee=provider_fee, + sats_to_usd=sats_to_usd, + ) + assert row["status"] == STATUS_OK + assert row["evidence"]["basis"] == "upstream_reported_usd" + + +# --- build_checklist -------------------------------------------------------- + + +class TestBuildChecklist: + def _row(self, row_id: str, status: str) -> dict[str, Any]: + return certification_row(row_id, status, "title", "detail") + + def test_all_ok_makes_all_goals_ok(self) -> None: + rows = [ + self._row("endpoint.reachable", STATUS_OK), + self._row("usage.capture", STATUS_OK), + self._row("cost.prompt_completion", STATUS_OK), + self._row("pricing.served_matches_configured", STATUS_OK), + self._row("pricing.enabled_models_served", STATUS_OK), + ] + checklist = build_checklist(rows) + assert len(checklist) == 4 + for item in checklist: + assert item["status"] == STATUS_OK + assert item["tick"] == "☑️" + + def test_one_fail_makes_its_goal_fail(self) -> None: + rows = [ + self._row("endpoint.reachable", STATUS_FAIL), + self._row("usage.capture", STATUS_OK), + self._row("cost.prompt_completion", STATUS_OK), + self._row("pricing.served_matches_configured", STATUS_OK), + self._row("pricing.enabled_models_served", STATUS_OK), + ] + checklist = build_checklist(rows) + goals = {item["goal"]: item["status"] for item in checklist} + assert goals["heartbeat"] == STATUS_FAIL + assert goals["usage_data"] == STATUS_OK + + def test_warn_makes_goal_warn(self) -> None: + rows = [ + self._row("endpoint.reachable", STATUS_OK), + self._row("usage.capture", STATUS_WARN), + self._row("cost.prompt_completion", STATUS_OK), + self._row("pricing.served_matches_configured", STATUS_OK), + self._row("pricing.enabled_models_served", STATUS_OK), + ] + checklist = build_checklist(rows) + goals = {item["goal"]: item["status"] for item in checklist} + assert goals["usage_data"] == STATUS_WARN + + def test_missing_row_makes_goal_warn(self) -> None: + rows = [self._row("endpoint.reachable", STATUS_OK)] + checklist = build_checklist(rows) + goals = {item["goal"]: item["status"] for item in checklist} + assert goals["heartbeat"] == STATUS_OK + assert goals["usage_data"] == STATUS_WARN # row absent → warn + + def test_pricing_goal_combines_two_rows(self) -> None: + """pricing_v1_models depends on TWO rows; one fail → goal fail.""" + rows = [ + self._row("endpoint.reachable", STATUS_OK), + self._row("usage.capture", STATUS_OK), + self._row("cost.prompt_completion", STATUS_OK), + self._row("pricing.served_matches_configured", STATUS_OK), + self._row("pricing.enabled_models_served", STATUS_FAIL), + ] + checklist = build_checklist(rows) + goals = {item["goal"]: item["status"] for item in checklist} + assert goals["pricing_v1_models"] == STATUS_FAIL + + def test_pricing_goal_ok_only_when_both_ok(self) -> None: + rows = [ + self._row("endpoint.reachable", STATUS_OK), + self._row("usage.capture", STATUS_OK), + self._row("cost.prompt_completion", STATUS_OK), + self._row("pricing.served_matches_configured", STATUS_OK), + self._row("pricing.enabled_models_served", STATUS_OK), + ] + checklist = build_checklist(rows) + goals = {item["goal"]: item["status"] for item in checklist} + assert goals["pricing_v1_models"] == STATUS_OK + + def test_fail_takes_precedence_over_warn(self) -> None: + rows = [ + self._row("endpoint.reachable", STATUS_FAIL), + self._row("usage.capture", STATUS_WARN), + self._row("cost.prompt_completion", STATUS_OK), + self._row("pricing.served_matches_configured", STATUS_OK), + self._row("pricing.enabled_models_served", STATUS_OK), + ] + checklist = build_checklist(rows) + goals = {item["goal"]: item["status"] for item in checklist} + assert goals["heartbeat"] == STATUS_FAIL + assert goals["usage_data"] == STATUS_WARN From 7f44389b597ee77cc2d48e4dbadbc1d6cfe85a7b Mon Sep 17 00:00:00 2001 From: 9qeklajc <211699015+9qeklajc@users.noreply.github.com> Date: Mon, 21 Sep 2026 15:53:14 +0000 Subject: [PATCH 04/18] fix(certification): harden against adversarial inputs found by tester agents MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- routstr/core/admin.py | 25 +- routstr/payment/usage.py | 19 +- routstr/upstream/certification.py | 322 ++++++++++++++--- tests/unit/test_certification_hardening.py | 388 +++++++++++++++++++++ 4 files changed, 691 insertions(+), 63 deletions(-) create mode 100644 tests/unit/test_certification_hardening.py 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 From 7ea551f407befafaaf0f39122f5e8ce2415e0cdb Mon Sep 17 00:00:00 2001 From: 9qeklajc <211699015+9qeklajc@users.noreply.github.com> Date: Mon, 21 Sep 2026 22:54:56 +0200 Subject: [PATCH 05/18] chore: declare respx as a dev dependency for the certify endpoint tests --- pyproject.toml | 1 + uv.lock | 14 ++++++++++++++ 2 files changed, 15 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 9f0eabd1..bb6f456f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,6 +36,7 @@ dev = [ "psutil>=5.9.0", "aiohttp>=3.9.0", "pytest-benchmark>=4.0.0", + "respx>=0.21", "routstr", ] diff --git a/uv.lock b/uv.lock index 6ca3e9ed..1da5b968 100644 --- a/uv.lock +++ b/uv.lock @@ -2584,6 +2584,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d7/8e/7540e8a2036f79a125c1d2ebadf69ed7901608859186c856fa0388ef4197/requests-2.33.1-py3-none-any.whl", hash = "sha256:4e6d1ef462f3626a1f0a0a9c42dd93c63bad33f9f1c1937509b8c5c8718ab56a", size = 64947, upload-time = "2026-03-30T16:09:13.83Z" }, ] +[[package]] +name = "respx" +version = "0.23.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "httpx", extra = ["socks"] }, +] +sdist = { url = "https://files.pythonhosted.org/packages/43/98/4e55c9c486404ec12373708d015ebce157966965a5ebe7f28ff2c784d41b/respx-0.23.1.tar.gz", hash = "sha256:242dcc6ce6b5b9bf621f5870c82a63997e8e82bc7c947f9ffe272b8f3dd5a780", size = 29243, upload-time = "2026-04-08T14:37:16.008Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/1d/4a/221da6ca167db45693d8d26c7dc79ccfc978a440251bf6721c9aaf251ac0/respx-0.23.1-py2.py3-none-any.whl", hash = "sha256:b18004b029935384bccfa6d7d9d74b4ec9af73a081cc28600fffc0447f4b8c1a", size = 25557, upload-time = "2026-04-08T14:37:14.613Z" }, +] + [[package]] name = "rich" version = "14.1.0" @@ -2645,6 +2657,7 @@ dev = [ { name = "pytest-asyncio" }, { name = "pytest-benchmark" }, { name = "pytest-cov" }, + { name = "respx" }, { name = "routstr" }, { name = "ruff" }, ] @@ -2680,6 +2693,7 @@ dev = [ { name = "pytest-asyncio", specifier = ">=0.24.0" }, { name = "pytest-benchmark", specifier = ">=4.0.0" }, { name = "pytest-cov", specifier = ">=6.1.1" }, + { name = "respx", specifier = ">=0.21" }, { name = "routstr", editable = "." }, { name = "ruff", specifier = ">=0.11.6" }, ] From 23a6eafbc641b52c1283e29a504d5fa7fcd2604a Mon Sep 17 00:00:00 2001 From: 9qeklajc <211699015+9qeklajc@users.noreply.github.com> Date: Mon, 21 Sep 2026 23:07:48 +0200 Subject: [PATCH 06/18] fix(tests): narrow provider.id to int for mypy in the certify endpoint tests --- tests/integration/test_certify_endpoint.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index 3f5b040f..49ec7de5 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -121,6 +121,7 @@ async def _seed_and_init( provider = await _make_provider( session, provider_fee=provider_fee, base_url=base_url ) + assert provider.id is not None session.add( _model_row( provider.id, @@ -475,6 +476,7 @@ async def test_certify_with_no_served_model( unusable pricing), the live checks should be skipped as warn, not crash.""" provider = await _make_provider(integration_session) + assert provider.id is not None # A negative prompt price makes has_usable_pricing() return False, # withholding the model from the served map. session_add = _model_row( From 0b8e07f834841d7254c73eeb3788d55f34a96271 Mon Sep 17 00:00:00 2001 From: 9qeklajc <211699015+9qeklajc@users.noreply.github.com> Date: Mon, 21 Sep 2026 23:24:06 +0200 Subject: [PATCH 07/18] refactor: simplify certification comments --- routstr/core/admin.py | 45 +++--- routstr/payment/usage.py | 10 +- routstr/upstream/certification.py | 170 ++++++--------------- tests/integration/test_certify_endpoint.py | 5 - tests/unit/test_certification_hardening.py | 41 ++--- 5 files changed, 76 insertions(+), 195 deletions(-) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index cb4c3af6..8a51d9c6 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1267,12 +1267,10 @@ def _served_model_for_provider(model_id: str, provider_pk: int) -> Model | None: class _ModelEvaluation: """One enabled model row's facts, built once and shared by every row. - ``configured`` is the fee-applied USD view built fresh from the row, or - ``None`` when the stored row could not be parsed (``build_error`` then - carries the exception). ``served`` is this provider's live candidate for - the model, or ``None`` when it is not being served at all — e.g. an - unusable stored price holds it back from the served map even though the - row itself is enabled. + ``configured`` is the fee-applied USD view of the row, or ``None`` when the + row could not be parsed (``build_error`` carries the exception). ``served`` + is this provider's live candidate, or ``None`` when the model is withheld + from the served map despite the row being enabled. """ model_id: str @@ -1422,10 +1420,9 @@ def _report_row_cache_rate( checked = 0 unknown: list[dict[str, object]] = [] for ev in evaluations: - # Scoped to served models only, like every sibling row: a model the - # routing algorithm withholds from the served map (e.g. an unusable - # stored price) has no cache-billing behaviour to certify here — - # ``pricing.enabled_models_served`` already flags it as unserved. + # Served models only, like every sibling row: an unserved model has no + # cache-billing behaviour to certify, and + # ``pricing.enabled_models_served`` already flags it. if ev.served is None or ev.configured is None: continue checked += 1 @@ -1514,18 +1511,15 @@ async def certify_upstream_provider( ) -> dict[str, object]: """Live certification checks for a configured upstream provider. - Unlike the read-only ``GET …/report``, this endpoint probes the - upstream over the network: it calls ``/models`` and sends a one-token - completion, then runs the node's own cost engine on the real response. - It never enters the billing path — no reservation, no Cashu, no wallet - — so it cannot spend the node's wallet. It costs at most one - completion's worth of upstream credit. + Unlike the read-only ``GET …/report``, this probes the upstream over the + network and runs the node's cost engine on the real response. It never + enters the billing path, so it costs at most one completion's worth of + upstream credit and nothing from the node's wallet. - The response carries the four ``pricing.*`` rows from the read-only - report (re-derived here so the certification is self-contained) plus - the five live/derived rows from - :mod:`routstr.upstream.certification`, and a ``checklist`` summarising - the four operator-facing goals with ``ok``/``warn``/``fail`` ticks. + Returns the read-only report's four ``pricing.*`` rows (re-derived here so + the certification is self-contained), the live rows from + :mod:`routstr.upstream.certification`, and a ``checklist`` of the + operator-facing goals. """ from ..payment.price import sats_usd_price from ..upstream.certification import ( @@ -1558,9 +1552,8 @@ async def certify_upstream_provider( model_id = payload.model_id if not model_id and enabled_rows: - # Pick the first enabled row that is actually being served — a - # model withheld from the served map would fail the chat probe for - # a reason unrelated to the endpoint's health. + # Prefer a served model: one withheld from the served map would fail + # the chat probe for a reason unrelated to the endpoint's health. for ev in evaluations: if ev.served is not None: model_id = ev.served.id @@ -1634,8 +1627,8 @@ async def certify_upstream_provider( ] else: sats_to_usd = sats_usd_price() - # Clamp the admin-supplied timeout: the probe must never be able to - # hold the request open indefinitely. + # Clamp the admin-supplied timeout so a probe cannot hold the request + # open indefinitely. requested = ( payload.timeout_seconds if payload.timeout_seconds is not None diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 7bee961b..d69ab55c 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -55,12 +55,10 @@ class NormalizedUsage(BaseModel): def parse_token_count(value: object) -> int: """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. + ``json.loads`` accepts bare ``Infinity``/``NaN`` and overflows ``1e999`` to + ``inf``, so an upstream can put them on the wire. ``int()`` raises on both, + which would turn a billing path into a 500; reject them like + ``is_usable_rate`` does instead. """ if isinstance(value, bool): return 0 diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index b36e9e46..a032eedb 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -1,33 +1,12 @@ -"""Certification checks for an upstream provider endpoint. +"""Live certification checks for an upstream provider endpoint. -PR #717 established the row contract — ``{id, status, title, detail, -evidence}`` with ``status`` in ``{ok, warn, fail}`` — and the four pricing -rows derived from the database row plus the in-process served map. Those -rows deliberately never touch the network. This module adds the checks that -*must* touch the network, and the checklist view that maps the -operator-facing goals onto rows: +Extends the read-only pricing rows, which never touch the network, with the +ones that must: a ``/models`` heartbeat and a one-token completion. -========================= ========================================= -Goal Row(s) -========================= ========================================= -Heartbeat ``endpoint.reachable`` -Usage data ``usage.capture`` -Cost data ``cost.prompt_completion`` -Pricing in ``/v1/models`` ``pricing.served_matches_configured``, - ``pricing.enabled_models_served`` -========================= ========================================= - -**Money safety.** Every live check calls the upstream directly with -``httpx`` — exactly like the existing ``POST /api/models/test`` probe — and -never enters the node's billing path. No reservation is taken, no Cashu -token is minted or spent, and the probe asks for a single token -(``max_tokens=1``). A probe therefore costs the operator at most one -completion's worth of upstream spend and nothing from the node's wallet. - -**Why a separate endpoint.** ``GET …/report`` promises the operator a -cheap, non-blocking read. A live probe can hang for the length of its -timeout and spends upstream credit, so it lives behind -``POST …/certify`` instead of being folded into the read. +Probes call the upstream directly with ``httpx``, never through the node's +billing path — no reservation, no Cashu, at most one token of upstream spend. +They sit behind ``POST …/certify`` rather than the read-only ``GET …/report`` +because they can block for the length of the timeout. """ from __future__ import annotations @@ -61,31 +40,22 @@ STATUS_FAIL = "fail" TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"} -# A probe must never be able to wedge an admin request. Fifteen seconds is -# generous for a `/models` listing or a one-token completion on a healthy -# upstream, and bounded enough that a dead host fails the row rather than -# the request. +# Bounded so a dead upstream fails the row rather than wedging 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. +# Ceiling for the caller-supplied timeout override. 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. +# The cheapest request that still exercises the usage/cost path. PROBE_MAX_TOKENS = 1 PROBE_PROMPT = "ping" -# The reservation ceiling is irrelevant to the token-priced path — it is -# only the amount held before settlement — but ``calculate_cost`` requires -# one. Any value at or above the real charge behaves identically. +# ``calculate_cost`` demands a reservation ceiling; any value at or above the +# real charge behaves identically. _PROBE_MAX_COST_MSATS = 1_000_000_000 -# Rounding in ``_calculate_from_tokens`` truncates the output component and -# folds the remainder into the input component, so a one-millisatoshi -# difference is arithmetic, not drift. +# ``_calculate_from_tokens`` truncates the output component and folds the +# remainder into the input one, so a one-msat difference is arithmetic. COST_TOLERANCE_MSATS = 1 @@ -96,12 +66,8 @@ def certification_row( detail: str, evidence: dict[str, Any] | None = None, ) -> dict[str, Any]: - """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. - """ + """Build one row, coercing ``evidence`` to a dict so the row contract + holds by construction rather than by caller discipline.""" return { "id": row_id, "status": status, @@ -116,13 +82,8 @@ def safe_row( 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. - """ + """Run a row builder, turning any raise into a ``fail`` row: the report is + the diagnostic, so it must never be the thing that 500s.""" try: return builder() except Exception as exc: # noqa: BLE001 - a raising check is a row status @@ -140,10 +101,8 @@ def safe_row( ) -# 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 -# it ``warn``. Kept as data so the checklist and the row set cannot drift. +# Operator-facing goals mapped onto the rows that decide them: ``ok`` only when +# every named row is ``ok``, ``fail`` if any fails, ``warn`` otherwise. CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = ( ( "heartbeat", @@ -169,7 +128,6 @@ CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = ( def build_checklist(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: - """Summarise the rows as the four operator-facing goals with ticks.""" by_id = {row["id"]: row for row in rows} checklist: list[dict[str, Any]] = [] for goal, label, row_ids in CHECKLIST_GOALS: @@ -295,25 +253,17 @@ async def probe_upstream( return result -# --------------------------------------------------------------------------- -# Row builders -# -# Every builder below is pure: it turns an already-fetched fact (a probe -# result, a model, a computed cost) into a row. The network lives only in -# ``probe_upstream`` and ``run_live_checks``, so a test can exercise each -# verdict — including the failure ones — without a socket. -# --------------------------------------------------------------------------- +# Row builders are pure: the network lives only in ``probe_upstream`` and +# ``run_live_checks``, so every verdict is testable without a socket. def endpoint_validity_row(base_url: str) -> dict[str, Any]: - """Check the configured base URL is a well-formed http(s) endpoint.""" parsed = urlparse(base_url or "") problems: list[str] = [] if parsed.scheme not in ("http", "https"): problems.append(f"scheme {parsed.scheme!r} is not http or https") - # ``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". + # ``netloc`` is truthy for a hostless authority like ``http://:8080``; + # only ``.hostname`` answers whether there is a host to connect to. if not parsed.hostname: problems.append("no host component") evidence: dict[str, Any] = { @@ -343,7 +293,6 @@ def endpoint_validity_row(base_url: str) -> dict[str, Any]: def heartbeat_row(probe: ProbeResult) -> dict[str, Any]: - """Check the upstream's ``/models`` responds — the heartbeat.""" evidence: dict[str, Any] = { "url": probe.models_url, "status_code": probe.models_status, @@ -377,7 +326,6 @@ 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.""" payload = probe.models_payload if not isinstance(payload, dict): return certification_row( @@ -402,8 +350,6 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]: }, ) - # An empty id is not an id — the CLI discovery path refuses it, so the - # row must not certify it either. ids = [ item["id"] for item in data @@ -436,10 +382,8 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]: def usage_capture_row(probe: ProbeResult) -> dict[str, Any]: """Check a completion comes back with token usage the node can bill on. - A missing ``usage`` object is the root of the ``(0+0)`` billing bug — - the node has nothing to price, so the request settles for free. That is - a real defect in the upstream's OpenAI compatibility, but it does not - make the endpoint unusable, so it is a ``warn`` rather than a ``fail``. + A missing ``usage`` object means the node has nothing to price and the + request settles for free. Broken, but still usable, so ``warn``. """ evidence: dict[str, Any] = { "url": probe.chat_url, @@ -532,14 +476,9 @@ def _truncate(value: Any, limit: int = 400) -> Any: def _reported_usd_cost(payload: dict[str, Any]) -> float: """The upstream-reported USD cost, or 0.0 when it reported none. - 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. 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. + Mirrors ``_resolve_usd_cost``'s priority and shares ``coerce_rate``, so + this helper and the engine agree on *whether* a cost was reported; only + the arithmetic below is re-derived independently. """ usage = payload.get("usage") if not isinstance(usage, dict): @@ -560,19 +499,12 @@ def _reported_usd_cost(payload: dict[str, Any]) -> float: def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int]: """Re-derive the token-priced charge independently of the engine. - ``_calculate_from_tokens`` prices at *msats per 1000 tokens*, rounds - each component to three decimals, ceilings the sum, then folds the - cache cost into the input component by truncating the output one. The - arithmetic is reproduced here — rather than calling the engine and - comparing it to itself — so a swapped input/output rate, a dropped - cache term or a changed rounding rule shows up as a mismatch. + Reproduces ``_calculate_from_tokens``'s arithmetic rather than calling the + engine and comparing it to itself, so a swapped rate, a dropped 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. + Returns ``(total_msats, input_msats, output_msats)``. Raises ``ValueError`` + on a non-finite rate, which would otherwise crash ``math.ceil`` downstream. """ input_rate = float(sats_pricing.prompt) * 1_000_000.0 output_rate = float(sats_pricing.completion) * 1_000_000.0 @@ -619,10 +551,8 @@ def cost_prompt_completion_row( ) -> dict[str, Any]: """Check the node's cost engine prices a real completion correctly. - Both the prompt and the completion component are checked: the engine - truncates the output component and folds the remainder into the input - component so that ``input + output == total`` exactly, which means a - wrong rate on *either* side shows up as a mismatch here. + Both components are checked, since the engine folds the truncated output + remainder into the input one to keep ``input + output == total``. """ from ..payment.cost_calculation import CostDataError @@ -842,17 +772,9 @@ async def run_live_checks( return rows -# --------------------------------------------------------------------------- -# Standalone runner -# -# ``certify_upstream_url`` deliberately reads nothing from the node's -# database: the point of the CLI is to certify a URL *before* it is -# configured, or one the operator does not want to write into the node at -# all. The four pricing rows therefore do not apply here — they compare a -# stored row against the served map, neither of which exists for a bare -# URL — and the cost row falls back to litellm's cost map (or explicit -# prices) instead of a configured row. -# --------------------------------------------------------------------------- +# The standalone runner certifies a URL before it is configured, so it reads +# nothing from the node's database: the pricing rows do not apply, and the cost +# row falls back to litellm's cost map or explicit prices. def _first_model_id(probe: ProbeResult) -> str | None: @@ -870,9 +792,8 @@ def _first_model_id(probe: ProbeResult) -> str | None: def _as_price(value: Any) -> float | None: """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. + Shares ``coerce_rate`` so an explicit ``--prompt-price`` is validated + exactly like a litellm-derived one. """ return coerce_rate(value) @@ -1036,7 +957,6 @@ async def certify_upstream_url( def render_checklist(result: dict[str, Any]) -> str: - """Render one certification result as the operator-facing checklist.""" target = result.get("target", {}) lines = [f"Upstream certification — {target.get('base_url')}"] if target.get("model_id"): @@ -1054,12 +974,8 @@ def render_checklist(result: dict[str, Any]) -> str: 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. - """ + """Move the app's stdout log handlers to stderr, so log records cannot + interleave with the report.""" import logging loggers = [logging.getLogger()] diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index 49ec7de5..69f432aa 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -224,7 +224,6 @@ async def test_certify_all_ok( assert "rows" in body assert "checklist" in body - # All live rows should be ok live_row_ids = [ "endpoint.validity", "endpoint.reachable", @@ -236,7 +235,6 @@ async def test_certify_all_ok( row = _find_row(body["rows"], row_id) assert row["status"] == "ok", f"{row_id}: {row}" - # All checklist goals should be ok for item in body["checklist"]: assert item["status"] == "ok", f"{item['goal']}: {item}" @@ -267,7 +265,6 @@ async def test_certify_heartbeat_fail_on_500( assert row["status"] == "fail" assert row["evidence"]["status_code"] == 500 - # heartbeat goal should be fail heartbeat_goal = next( item for item in body["checklist"] if item["goal"] == "heartbeat" ) @@ -338,7 +335,6 @@ async def test_certify_usage_warn_when_no_usage( respx.get("https://certify-upstream.example/v1/models").mock( return_value=Response(200, json=_mock_models_response()) ) - # No "usage" key in the chat response respx.post("https://certify-upstream.example/v1/chat/completions").mock( return_value=Response( 200, @@ -494,7 +490,6 @@ async def test_certify_with_no_served_model( ) assert resp.status_code == 200, resp.text body = resp.json() - # Live rows should be warn (skipped) for row_id in ["endpoint.reachable", "usage.capture", "cost.prompt_completion"]: row = _find_row(body["rows"], row_id) assert row["status"] == "warn", f"{row_id}: {row}" diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py index ae707448..52f2fb02 100644 --- a/tests/unit/test_certification_hardening.py +++ b/tests/unit/test_certification_hardening.py @@ -50,13 +50,10 @@ def _probe(**kwargs: Any) -> ProbeResult: ) -# --------------------------------------------------------------------------- -# Defect: a non-finite token count crashed the billing path. -# +# Regression: 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: @@ -111,10 +108,8 @@ class TestNonFiniteTokenCounts: assert row["status"] == STATUS_WARN -# --------------------------------------------------------------------------- -# Defect: ``certification_row`` stored non-dict evidence verbatim, so the +# Regression: ``certification_row`` stored non-dict evidence verbatim, so the # row contract ("evidence is always a dict") held only by caller discipline. -# --------------------------------------------------------------------------- class TestEvidenceContract: @@ -128,10 +123,8 @@ class TestEvidenceContract: assert row["evidence"] == {"a": 1} -# --------------------------------------------------------------------------- -# Defect: ``http://:8080/v1`` was certified as a valid endpoint because +# Regression: ``http://:8080/v1`` was certified as a valid endpoint because # ``netloc`` is truthy for a hostless authority. -# --------------------------------------------------------------------------- class TestEndpointValidity: @@ -157,11 +150,9 @@ class TestEndpointValidity: assert row["status"] == STATUS_OK, url -# --------------------------------------------------------------------------- -# Defect: the payload builders called ``.get()`` on whatever they were +# Regression: the payload builders called ``.get()`` on whatever they were # given, so a wrong-typed body raised AttributeError instead of producing a # verdict. -# --------------------------------------------------------------------------- class TestPayloadTypeGuards: @@ -184,10 +175,8 @@ class TestPayloadTypeGuards: assert row["evidence"]["usable_ids"] == 0 -# --------------------------------------------------------------------------- -# Defect: an empty-string id was counted as "usable" by the payload row but +# Regression: 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: @@ -204,10 +193,8 @@ class TestModelIdAgreement: assert row["evidence"]["usable_ids"] == 1 -# --------------------------------------------------------------------------- -# Defect: the independent cost re-derivation disagreed with the engine on +# Regression: the independent cost re-derivation disagreed with the engine on # coercion (numeric strings, booleans), manufacturing false failures. -# --------------------------------------------------------------------------- class TestReportedCostCoercionParity: @@ -230,10 +217,8 @@ class TestReportedCostCoercionParity: assert _reported_usd_cost(payload) == pytest.approx(0.001) -# --------------------------------------------------------------------------- -# Defect: ``_expected_token_msats`` ran ``math.ceil`` on a non-finite sum, +# Regression: ``_expected_token_msats`` ran ``math.ceil`` on a non-finite sum, # raising an opaque error instead of a describable one. -# --------------------------------------------------------------------------- class TestNonFinitePricing: @@ -267,9 +252,7 @@ class TestNonFinitePricing: assert total == inp + outp -# --------------------------------------------------------------------------- -# Defect: a row builder raising escaped as a 500 from the admin endpoint. -# --------------------------------------------------------------------------- +# Regression: a row builder raising escaped as a 500 from the admin endpoint. class TestSafeRow: @@ -289,10 +272,8 @@ class TestSafeRow: assert row["status"] == STATUS_OK -# --------------------------------------------------------------------------- -# Defect: explicit ``--prompt-price`` bypassed validation, so a negative +# Regression: explicit ``--prompt-price`` bypassed validation, so a negative # rate could be fed into the cost engine. -# --------------------------------------------------------------------------- class TestExplicitPriceValidation: @@ -314,12 +295,10 @@ class TestExplicitPriceValidation: assert _as_price("1e-7") == pytest.approx(1e-7) -# --------------------------------------------------------------------------- -# Defect: the standalone CLI was dead on arrival — ``sats_usd_price()`` +# Regression: 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]: From 721eb4ff8d050ef00270422fb169177454c915d7 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 21:19:36 +0200 Subject: [PATCH 08/18] feat: add model path certification --- routstr/core/admin.py | 94 ++- routstr/upstream/certification.py | 70 +- routstr/upstream/certification_cache.py | 586 +++++++++++++++ routstr/upstream/model_paths.py | 52 ++ tests/integration/test_certify_endpoint.py | 411 ++++++++++- tests/unit/test_certification.py | 5 +- tests/unit/test_certification_cache.py | 499 +++++++++++++ tests/unit/test_certification_hardening.py | 4 +- ui/components/provider-card.tsx | 20 + .../provider-certification-dialog.tsx | 697 ++++++++++++++++++ ui/lib/api/services/admin.ts | 59 ++ 11 files changed, 2489 insertions(+), 8 deletions(-) create mode 100644 routstr/upstream/certification_cache.py create mode 100644 tests/unit/test_certification_cache.py create mode 100644 ui/components/provider-certification-dialog.tsx diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 8a51d9c6..bd5b4e18 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -28,6 +28,7 @@ from .db import ( CashuTransaction, CliToken, LightningInvoice, + ModelPathRow, ModelRow, UpstreamProviderRow, create_session, @@ -1234,6 +1235,30 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: m for m in upstream_models if m.id not in db_model_ids ] + path_result = await session.exec( + select(ModelPathRow).where( + ModelPathRow.upstream_provider_id == provider_pk + ) + ) + path_rows = list(path_result.all()) + paths_by_public_id: dict[str, list[dict[str, object]]] = {} + for row in path_rows: + paths_by_public_id.setdefault(row.model_id.lower(), []).append( + { + "path": row.path, + "endpoint_tag": row.endpoint_tag, + "endpoint_name": row.endpoint_name, + } + ) + + from ..upstream.model_paths import public_model_id + + certification_paths: dict[str, list[dict[str, object]]] = {} + for model in [*db_models, *filtered_remote_models]: + forwarded_id = model.forwarded_model_id or model.id + paths = paths_by_public_id.get(public_model_id(forwarded_id).lower(), []) + certification_paths[model.id] = paths + return { "provider": { "id": provider.id, @@ -1247,6 +1272,7 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: # missing one; show the operator the value that needs fixing. "db_models": [json_compliant(m.dict()) for m in db_models], "remote_models": [json_compliant(m.dict()) for m in filtered_remote_models], + "certification_paths": certification_paths, } @@ -1499,7 +1525,9 @@ async def get_upstream_provider_report(provider_id: str) -> dict[str, object]: class CertifyRequest(BaseModel): model_id: str | None = None + model_path: str | None = None timeout_seconds: float | None = None + check_cache: bool = True @admin_router.post( @@ -1540,6 +1568,36 @@ async def certify_upstream_provider( ) enabled_rows = list(result.all()) + endpoint_tag: str | None = None + selected_path: ModelPathRow | None = None + if payload.model_path is not None: + from ..proxy import _model_ids_match + from ..upstream.model_paths import decode_model_path + + selector = decode_model_path(payload.model_path) + if selector is None: + raise HTTPException(status_code=400, detail="Malformed model path") + if payload.model_id is None or not _model_ids_match( + payload.model_id, selector.model_id + ): + raise HTTPException( + status_code=400, + detail="Model path does not match the selected model", + ) + path_result = await session.exec( + select(ModelPathRow).where( + ModelPathRow.upstream_provider_id == provider_pk, + ModelPathRow.path == payload.model_path, + ) + ) + selected_path = path_result.first() + if selected_path is None: + raise HTTPException( + status_code=400, + detail="Model path is not available for this provider", + ) + endpoint_tag = selector.endpoint_tag + evaluations = [ _evaluate_model_row(row, provider, provider_pk) for row in enabled_rows ] @@ -1561,7 +1619,7 @@ async def certify_upstream_provider( if model_id is None: model_id = enabled_rows[0].id - from ..proxy import get_candidates + from ..proxy import get_candidates, get_upstreams model_obj = None if model_id: @@ -1581,11 +1639,33 @@ async def certify_upstream_provider( if model.upstream_provider_id == provider_pk: model_obj = model break + + # A model selected from the provider's discovered catalog may not have + # a database override and therefore may not appear in get_candidates(). + # The active upstream cache carries the same fee-adjusted USD and sats + # pricing used by the proxy, so it is the authoritative fallback for a + # pre-configuration certification probe. + if model_obj is None: + for upstream in get_upstreams(): + if getattr(upstream, "db_id", None) != provider_pk: + continue + model_obj = next( + ( + model + for model in upstream.get_cached_models() + if model.id == model_id + or model.forwarded_model_id == model_id + ), + None, + ) + if model_obj is not None: + break if model_obj is None: from ..upstream.certification import ( STATUS_WARN, certification_row, ) + from ..upstream.certification_cache import skipped_cache_rows live_rows = [ certification_row( @@ -1624,9 +1704,19 @@ async def certify_upstream_provider( "Skipped — no model to probe.", {}, ), + *skipped_cache_rows("Skipped — no model to probe."), ] else: sats_to_usd = sats_usd_price() + if selected_path is not None: + from ..upstream.model_paths import apply_model_path_pricing + + model_obj = apply_model_path_pricing( + model_obj, + selected_path, + provider.provider_fee, + sats_to_usd, + ) # Clamp the admin-supplied timeout so a probe cannot hold the request # open indefinitely. requested = ( @@ -1642,6 +1732,8 @@ async def certify_upstream_provider( provider_fee=provider.provider_fee, sats_to_usd=sats_to_usd, timeout=timeout, + check_cache=payload.check_cache, + endpoint_tag=endpoint_tag, ) rows = pricing_rows + live_rows diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index a032eedb..5a87c8d0 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -28,6 +28,7 @@ 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 +from .model_paths import is_openrouter_base_url if TYPE_CHECKING: from ..payment.models import Model @@ -124,6 +125,16 @@ CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = ( "Pricing in /v1/models — cost updates reflected in the models list", ("pricing.served_matches_configured", "pricing.enabled_models_served"), ), + ( + "caching", + "Prompt caching — cache hits reported and billed at the cache rate", + ("cache.reported", "cache.billing"), + ), + ( + "margin", + "Margin — node charge covers the upstream's cost", + ("cost.margin",), + ), ) @@ -159,6 +170,7 @@ class ProbeResult: base_url: str models_url: str chat_url: str + endpoint_tag: str | None = None models_status: int | None = None models_payload: dict[str, Any] | None = None models_error: str | None = None @@ -174,6 +186,7 @@ async def probe_upstream( api_key: str, model_id: str, *, + endpoint_tag: str | None = None, client: httpx.AsyncClient | None = None, timeout: float = PROBE_TIMEOUT_SECONDS, ) -> ProbeResult: @@ -187,6 +200,7 @@ async def probe_upstream( base_url=base_url, models_url=f"{base}/models", chat_url=f"{base}/chat/completions", + endpoint_tag=endpoint_tag, ) headers = {"Content-Type": "application/json"} if api_key: @@ -226,6 +240,11 @@ async def probe_upstream( "max_tokens": PROBE_MAX_TOKENS, "stream": False, } + if endpoint_tag: + request_body["provider"] = { + "order": [endpoint_tag], + "allow_fallbacks": False, + } try: response = await client.post( result.chat_url, json=request_body, headers=headers @@ -702,12 +721,19 @@ async def run_live_checks( client: httpx.AsyncClient | None = None, timeout: float = PROBE_TIMEOUT_SECONDS, pricing_known: bool = True, + check_cache: bool = True, + endpoint_tag: str | None = None, ) -> list[dict[str, Any]]: - """Probe one upstream once and build the five live/derived rows.""" + """Probe one upstream and build the live/derived rows. + + ``check_cache`` adds the prompt-cache and margin rows, which cost two + more completions against a long prompt. + """ probe = await probe_upstream( base_url, api_key, model.forwarded_model_id or model.id, + endpoint_tag=endpoint_tag, client=client, timeout=timeout, ) @@ -769,6 +795,37 @@ async def run_live_checks( ), ) ) + + from .certification_cache import run_cache_checks, skipped_cache_rows + + if not check_cache: + rows.extend(skipped_cache_rows("Skipped — cache checks disabled.")) + elif probe.chat_payload is None: + rows.extend( + skipped_cache_rows("Skipped — the completion probe did not succeed.") + ) + elif is_openrouter_base_url(base_url) and endpoint_tag is None: + rows.extend( + skipped_cache_rows( + "Skipped — select an exact OpenRouter model path so both cache " + "requests use the same upstream endpoint." + ) + ) + else: + rows.extend( + await run_cache_checks( + base_url, + api_key, + model, + provider_fee=provider_fee, + sats_to_usd=sats_to_usd, + probe_payload=probe.chat_payload, + client=client, + timeout=timeout, + pricing_known=pricing_known, + endpoint_tag=endpoint_tag, + ) + ) return rows @@ -877,6 +934,7 @@ async def certify_upstream_url( timeout: float = PROBE_TIMEOUT_SECONDS, sats_usd_price: float | None = None, client: httpx.AsyncClient | None = None, + check_cache: bool = True, ) -> dict[str, Any]: """Certify an arbitrary upstream URL without touching the node's DB.""" from ..payment.models import litellm_cost_entry @@ -911,6 +969,9 @@ async def certify_upstream_url( {}, ), ] + from .certification_cache import skipped_cache_rows + + rows.extend(skipped_cache_rows("Skipped — no model to probe.")) return { "target": target, "rows": rows, @@ -952,6 +1013,7 @@ async def certify_upstream_url( client=client, timeout=timeout, pricing_known=pricing_known, + check_cache=check_cache, ) return {"target": target, "rows": rows, "checklist": build_checklist(rows)} @@ -1061,6 +1123,11 @@ def main(argv: list[str] | None = None) -> int: "parseable — use this in pipelines." ), ) + parser.add_argument( + "--no-cache", + action="store_true", + help="Skip the prompt-cache and margin rows (saves two long completions)", + ) args = parser.parse_args(argv) async def _run_all() -> list[dict[str, Any]]: @@ -1076,6 +1143,7 @@ def main(argv: list[str] | None = None) -> int: provider_fee=args.provider_fee, timeout=args.timeout, sats_usd_price=args.sats_usd_price, + check_cache=not args.no_cache, ) ) return results diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py new file mode 100644 index 00000000..3ee74e7a --- /dev/null +++ b/routstr/upstream/certification_cache.py @@ -0,0 +1,586 @@ +"""Prompt-cache and margin certification for an upstream provider. + +Three questions the one-token probe cannot answer: + +* does the upstream *report* prompt-cache hits in a dialect the node parses, +* does the node bill cached reads at the discounted rate (client side), and +* does the node's charge cover what the upstream charged (node side). + +The cache probe sends the same long system prompt twice; the second call is +the one expected to report cached reads. Calls go straight to the upstream +with ``httpx`` and never enter the billing path. +""" + +from __future__ import annotations + +import time +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any + +import httpx + +from ..core.logging import get_logger +from ..payment.cost_calculation import CostDataError, calculate_cost +from ..payment.usage import NormalizedUsage, normalize_usage +from .certification import ( + _PROBE_MAX_COST_MSATS, + COST_TOLERANCE_MSATS, + PROBE_MAX_TOKENS, + PROBE_TIMEOUT_SECONDS, + STATUS_FAIL, + STATUS_OK, + STATUS_WARN, + _expected_token_msats, + _expected_usd_msats, + _reported_usd_cost, + certification_row, + safe_row, +) + +if TYPE_CHECKING: + from ..payment.models import Model + +logger = get_logger(__name__) + +# OpenAI caches prefixes of 1024+ tokens; Anthropic Haiku needs 2048+. The +# filler lands around 3000 tokens so every dialect can hit its threshold. +CACHE_PROBE_LINES = 220 +CACHE_PROBE_QUESTION = "Reply with the single word: ok" + +ROW_REPORTED = "cache.reported" +ROW_BILLING = "cache.billing" +ROW_MARGIN = "cost.margin" + +TITLE_REPORTED = "Upstream reports prompt-cache hits" +TITLE_BILLING = "Cached tokens billed at the cache-read rate" +TITLE_MARGIN = "Node charge covers upstream cost" + + +def cache_probe_prefix() -> str: + lines = [ + "You are a certification probe. Ignore the reference table below and " + "answer the final question with one word." + ] + for index in range(CACHE_PROBE_LINES): + lines.append( + f"Reference row {index:04d}: token {index * 7919 % 10007} maps to " + f"slot {index * 104729 % 1009} in region {index % 17}." + ) + return "\n".join(lines) + + +@dataclass +class CacheProbeResult: + """Two identical completions; the second should read from the cache.""" + + chat_url: str + request_format: str = "cache_control" + endpoint_tag: str | None = None + statuses: list[int | None] = field(default_factory=list) + payloads: list[dict[str, Any] | None] = field(default_factory=list) + errors: list[str | None] = field(default_factory=list) + latencies_ms: list[float | None] = field(default_factory=list) + + @property + def second_payload(self) -> dict[str, Any] | None: + return self.payloads[1] if len(self.payloads) > 1 else None + + @property + def second_error(self) -> str | None: + if len(self.errors) > 1: + return self.errors[1] + return self.errors[0] if self.errors else "cache probe did not run" + + +def _request_body( + model_id: str, prefix: str, fmt: str, endpoint_tag: str | None +) -> dict[str, Any]: + system: Any + if fmt == "cache_control": + system = [ + { + "type": "text", + "text": prefix, + "cache_control": {"type": "ephemeral"}, + } + ] + else: + system = prefix + body: dict[str, Any] = { + "model": model_id, + "messages": [ + {"role": "system", "content": system}, + {"role": "user", "content": CACHE_PROBE_QUESTION}, + ], + "max_tokens": PROBE_MAX_TOKENS, + "stream": False, + } + if endpoint_tag: + body["provider"] = { + "order": [endpoint_tag], + "allow_fallbacks": False, + } + return body + + +async def _post_completion( + client: httpx.AsyncClient, + url: str, + body: dict[str, Any], + headers: dict[str, str], +) -> tuple[int | None, dict[str, Any] | None, str | None, float]: + started = time.monotonic() + try: + response = await client.post(url, json=body, headers=headers) + except Exception as exc: # noqa: BLE001 - transport failure is a row status + latency = round((time.monotonic() - started) * 1000, 2) + return None, None, f"{type(exc).__name__}: {exc}", latency + latency = round((time.monotonic() - started) * 1000, 2) + try: + payload = response.json() + except Exception as exc: # noqa: BLE001 - any decode failure is the signal + return response.status_code, None, f"{type(exc).__name__}: {exc}", latency + if not isinstance(payload, dict): + return ( + response.status_code, + None, + f"expected a JSON object, got {type(payload).__name__}", + latency, + ) + return response.status_code, payload, None, latency + + +def _record( + result: CacheProbeResult, + outcome: tuple[int | None, dict[str, Any] | None, str | None, float], +) -> None: + status, payload, error, latency = outcome + result.statuses.append(status) + result.payloads.append(payload) + result.errors.append(error) + result.latencies_ms.append(latency) + + +def _is_2xx(status: int | None) -> bool: + return status is not None and 200 <= status < 300 + + +async def probe_cache( + base_url: str, + api_key: str, + model_id: str, + *, + endpoint_tag: str | None = None, + client: httpx.AsyncClient | None = None, + timeout: float = PROBE_TIMEOUT_SECONDS, +) -> CacheProbeResult: + """Send the same long prompt twice. + + The first attempt marks the prefix with an Anthropic-style + ``cache_control`` part. Upstreams that reject the part get a plain string + retry, and the second call mirrors whichever format succeeded. + """ + base = base_url.rstrip("/") + result = CacheProbeResult( + chat_url=f"{base}/chat/completions", endpoint_tag=endpoint_tag + ) + headers = {"Content-Type": "application/json"} + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + prefix = cache_probe_prefix() + + owns_client = client is None + if client is None: + client = httpx.AsyncClient(timeout=timeout) + try: + first = await _post_completion( + client, + result.chat_url, + _request_body(model_id, prefix, "cache_control", endpoint_tag), + headers, + ) + if not _is_2xx(first[0]) and first[0] is not None: + result.request_format = "plain" + first = await _post_completion( + client, + result.chat_url, + _request_body(model_id, prefix, "plain", endpoint_tag), + headers, + ) + _record(result, first) + if not _is_2xx(first[0]): + return result + second = await _post_completion( + client, + result.chat_url, + _request_body(model_id, prefix, result.request_format, endpoint_tag), + headers, + ) + _record(result, second) + finally: + if owns_client: + await client.aclose() + return result + + +def _raw_cache_keys(value: Any, path: str = "") -> list[str]: + """Paths of positive numeric fields whose name mentions a cache.""" + found: list[str] = [] + if isinstance(value, dict): + for key, item in value.items(): + child = f"{path}.{key}" if path else str(key) + if "cach" in str(key).lower() and isinstance(item, (int, float)): + if not isinstance(item, bool) and item > 0: + found.append(child) + found.extend(_raw_cache_keys(item, child)) + return found + + +def _usage_of(payload: dict[str, Any] | None) -> NormalizedUsage | None: + if not isinstance(payload, dict): + return None + try: + return normalize_usage(payload.get("usage")) + except Exception: # noqa: BLE001 - a malformed usage object is a row status + return None + + +def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]: + evidence: dict[str, Any] = { + "url": probe.chat_url, + "request_format": probe.request_format, + "endpoint_tag": probe.endpoint_tag, + "statuses": probe.statuses, + "latencies_ms": probe.latencies_ms, + } + payload = probe.second_payload + if payload is None or not _is_2xx(probe.statuses[-1] if probe.statuses else None): + evidence["error"] = probe.second_error + return certification_row( + ROW_REPORTED, + STATUS_FAIL, + TITLE_REPORTED, + f"The cache probe did not get two successful completions: " + f"{probe.second_error or 'no response body'}.", + evidence, + ) + + first_usage = _usage_of(probe.payloads[0]) + second_usage = _usage_of(payload) + evidence["first_usage"] = first_usage.dict() if first_usage else None + evidence["second_usage"] = second_usage.dict() if second_usage else None + + if second_usage is not None and second_usage.cache_read_tokens > 0: + return certification_row( + ROW_REPORTED, + STATUS_OK, + TITLE_REPORTED, + f"The repeated prompt reported {second_usage.cache_read_tokens} " + f"cached tokens (first call wrote {first_usage.cache_write_tokens if first_usage else 0}).", + evidence, + ) + + raw_keys = _raw_cache_keys(payload.get("usage")) + if raw_keys: + evidence["unrecognised_cache_fields"] = raw_keys + return certification_row( + ROW_REPORTED, + STATUS_FAIL, + TITLE_REPORTED, + "The upstream reported cache tokens under fields the node does not " + f"parse ({', '.join(raw_keys)}); cached reads would be billed at " + "the full input rate.", + evidence, + ) + return certification_row( + ROW_REPORTED, + STATUS_WARN, + TITLE_REPORTED, + "Two identical prompts produced no cache hit. Either the model does " + "not support prompt caching or the upstream hides it; clients pay the " + "full input rate on repeated prompts.", + evidence, + ) + + +def cache_billing_row( + *, + model: Model, + probe: CacheProbeResult, + cost_data: Any, + pricing_known: bool = True, +) -> dict[str, Any]: + usage = _usage_of(probe.second_payload) + evidence: dict[str, Any] = {"model_id": model.id} + if usage is None or usage.cache_read_tokens <= 0: + return certification_row( + ROW_BILLING, + STATUS_WARN, + TITLE_BILLING, + "No cached reads were reported, so there is nothing to price.", + evidence, + ) + if model.sats_pricing is None or not pricing_known: + return certification_row( + ROW_BILLING, + STATUS_WARN, + TITLE_BILLING, + "No pricing is known for this model, so the cache discount cannot " + "be verified.", + evidence, + ) + if isinstance(cost_data, CostDataError): + evidence["error"] = cost_data.message + return certification_row( + ROW_BILLING, + STATUS_FAIL, + TITLE_BILLING, + f"The cost engine could not price the cached completion: " + f"{cost_data.message}.", + evidence, + ) + + pricing = model.sats_pricing + cache_read_rate = float(pricing.input_cache_read or 0.0) + input_rate = float(pricing.prompt) + full_usage = NormalizedUsage( + input_tokens=usage.input_tokens + + usage.cache_read_tokens + + usage.cache_write_tokens, + output_tokens=usage.output_tokens, + ) + try: + expected_total, _, _ = _expected_token_msats(pricing, usage) + full_total, _, _ = _expected_token_msats(pricing, full_usage) + except (ValueError, OverflowError) as exc: + evidence["error"] = f"{type(exc).__name__}: {exc}" + return certification_row( + ROW_BILLING, + STATUS_FAIL, + TITLE_BILLING, + f"The expected charge could not be derived: {exc}.", + evidence, + ) + + actual_total = int(cost_data.total_msats) + reported_usd = _reported_usd_cost(probe.second_payload or {}) + evidence.update( + { + "usage": usage.dict(), + "cache_read_rate_sats": cache_read_rate, + "input_rate_sats": input_rate, + "actual_total_msats": actual_total, + "expected_total_msats": expected_total, + "full_price_total_msats": full_total, + "reported_usd": reported_usd or None, + } + ) + + if reported_usd > 0: + return certification_row( + ROW_BILLING, + STATUS_OK, + TITLE_BILLING, + f"Billed {actual_total} msats from the upstream-reported cost, " + f"which already carries the cache discount " + f"(full token price would be {full_total} msats).", + evidence, + ) + if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS: + return certification_row( + ROW_BILLING, + STATUS_FAIL, + TITLE_BILLING, + f"The engine charged {actual_total} msats but the configured cache " + f"rate implies {expected_total} msats.", + evidence, + ) + if cache_read_rate <= 0.0 or cache_read_rate >= input_rate: + return certification_row( + ROW_BILLING, + STATUS_WARN, + TITLE_BILLING, + f"Cached reads are billed at the full input rate ({actual_total} " + "msats) because no discounted cache-read rate is configured; " + "clients pay more than the upstream charges.", + evidence, + ) + return certification_row( + ROW_BILLING, + STATUS_OK, + TITLE_BILLING, + f"Charged {actual_total} msats for {usage.cache_read_tokens} cached " + f"tokens, {full_total - actual_total} msats below the full input price.", + evidence, + ) + + +def cost_margin_row( + *, + model: Model, + payloads: list[dict[str, Any] | None], + provider_fee: float, + sats_to_usd: float, + pricing_known: bool = True, +) -> dict[str, Any]: + """Configured token pricing must cover what the upstream reports charging. + + Responses that carry a USD cost are billed from it, so they cannot lose + money themselves; they are used here as a price sample. The configured + token pricing is what every other path bills from (streams, upstreams + that omit cost, the served ``/v1/models`` list), so a sample where it + falls below the fee-adjusted upstream cost means those paths underprice. + Upstreams that report no cost give no sample and the row stays a warn. + """ + evidence: dict[str, Any] = { + "model_id": model.id, + "provider_fee": provider_fee, + "sats_usd_price": sats_to_usd, + "samples": [], + } + if model.sats_pricing is None or not pricing_known: + return certification_row( + ROW_MARGIN, + STATUS_WARN, + TITLE_MARGIN, + "No pricing is known for this model, so the margin cannot be verified.", + evidence, + ) + + samples: list[dict[str, Any]] = [] + short: list[str] = [] + for payload in payloads: + if not isinstance(payload, dict): + continue + reported_usd = _reported_usd_cost(payload) + usage = _usage_of(payload) + if reported_usd <= 0 or usage is None: + continue + try: + configured_total, _, _ = _expected_token_msats(model.sats_pricing, usage) + upstream_total = _expected_usd_msats( + reported_usd, provider_fee, sats_to_usd + ) + except (ValueError, OverflowError) as exc: + evidence["error"] = f"{type(exc).__name__}: {exc}" + return certification_row( + ROW_MARGIN, + STATUS_FAIL, + TITLE_MARGIN, + f"The margin could not be derived: {exc}.", + evidence, + ) + samples.append( + { + "usage": usage.dict(), + "reported_usd": reported_usd, + "upstream_msats_with_fee": upstream_total, + "configured_msats": configured_total, + } + ) + if configured_total + COST_TOLERANCE_MSATS < upstream_total: + short.append(f"{configured_total} < {upstream_total}") + evidence["samples"] = samples + + if not samples: + return certification_row( + ROW_MARGIN, + STATUS_WARN, + TITLE_MARGIN, + "The upstream does not report a cost, so the margin cannot be " + "verified live. Keep configured prices at or above the upstream's " + "list price.", + evidence, + ) + if short: + return certification_row( + ROW_MARGIN, + STATUS_FAIL, + TITLE_MARGIN, + "Configured pricing is below the upstream's reported cost " + f"(configured < upstream msats: {'; '.join(short)}); token-billed " + "requests lose money.", + evidence, + ) + return certification_row( + ROW_MARGIN, + STATUS_OK, + TITLE_MARGIN, + f"Configured pricing covers the upstream's reported cost on " + f"{len(samples)} sampled completion(s).", + evidence, + ) + + +async def _price_payload( + payload: dict[str, Any] | None, model: Model, provider_fee: float +) -> Any: + if payload is None: + return CostDataError( + message="the cache probe did not succeed", code="no_completion" + ) + try: + return await calculate_cost( + payload, _PROBE_MAX_COST_MSATS, model_obj=model, provider_fee=provider_fee + ) + except Exception as exc: # noqa: BLE001 - a raising engine is a fail row + return CostDataError( + message=f"{type(exc).__name__}: {exc}", code="pricing_error" + ) + + +async def run_cache_checks( + base_url: str, + api_key: str, + model: Model, + *, + provider_fee: float, + sats_to_usd: float, + probe_payload: dict[str, Any] | None, + client: httpx.AsyncClient | None = None, + timeout: float = PROBE_TIMEOUT_SECONDS, + pricing_known: bool = True, + endpoint_tag: str | None = None, +) -> list[dict[str, Any]]: + """Run the cache probe and build the three cache/margin rows.""" + probe = await probe_cache( + base_url, + api_key, + model.forwarded_model_id or model.id, + endpoint_tag=endpoint_tag, + client=client, + timeout=timeout, + ) + cost_data = await _price_payload(probe.second_payload, model, provider_fee) + return [ + safe_row(ROW_REPORTED, TITLE_REPORTED, lambda: cache_reported_row(probe)), + safe_row( + ROW_BILLING, + TITLE_BILLING, + lambda: cache_billing_row( + model=model, + probe=probe, + cost_data=cost_data, + pricing_known=pricing_known, + ), + ), + safe_row( + ROW_MARGIN, + TITLE_MARGIN, + lambda: cost_margin_row( + model=model, + payloads=[probe_payload, *probe.payloads], + provider_fee=provider_fee, + sats_to_usd=sats_to_usd, + pricing_known=pricing_known, + ), + ), + ] + + +def skipped_cache_rows(reason: str) -> list[dict[str, Any]]: + return [ + certification_row(ROW_REPORTED, STATUS_WARN, TITLE_REPORTED, reason, {}), + certification_row(ROW_BILLING, STATUS_WARN, TITLE_BILLING, reason, {}), + certification_row(ROW_MARGIN, STATUS_WARN, TITLE_MARGIN, reason, {}), + ] diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index fedb042e..0ac4e09d 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -37,6 +37,7 @@ from ..core.logging import get_logger if TYPE_CHECKING: from sqlmodel.ext.asyncio.session import AsyncSession + from ..payment.models import Model from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -884,6 +885,57 @@ def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None: model["sats_pricing"] = priced.sats_pricing.dict() +def apply_model_path_pricing( + model: "Model", + row: ModelPathRow, + provider_fee: float, + sats_to_usd: float, +) -> "Model": + """Return ``model`` priced from an exact endpoint path's own rates. + + Direct paths already use the provider model cache and therefore carry the + same pricing as ``model``. OpenRouter endpoint rows instead contain raw, + endpoint-specific USD rates; certification must use those rates when its + requests are pinned to that endpoint. + """ + if row.endpoint_tag is None: + return model + + from ..payment.models import ( + Pricing, + _calculate_usd_max_costs, + _update_model_sats_pricing, + backfill_cache_pricing, + ) + + try: + metadata = json.loads(row.model_metadata) + if not isinstance(metadata, dict) or not isinstance( + metadata.get("pricing"), dict + ): + return model + pricing = backfill_cache_pricing( + model.forwarded_model_id or row.model_id, + Pricing.parse_obj(metadata["pricing"]), + ) + pricing = Pricing.parse_obj( + {key: float(value) * provider_fee for key, value in pricing.dict().items()} + ) + priced = model.copy(update={"pricing": pricing, "sats_pricing": None}) + ( + pricing.max_prompt_cost, + pricing.max_completion_cost, + pricing.max_cost, + ) = _calculate_usd_max_costs(priced) + return _update_model_sats_pricing(priced, sats_to_usd) + except Exception as exc: + logger.warning( + "Could not apply model-path pricing for certification", + extra={"model_id": model.id, "path": row.path, "error": str(exc)}, + ) + return model + + def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]: endpoint = None if row.endpoint_tag or row.endpoint_name: diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index 69f432aa..c7aa9151 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -20,8 +20,9 @@ from httpx import AsyncClient, Response from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.admin import admin_sessions -from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.core.db import ModelPathRow, ModelRow, UpstreamProviderRow from routstr.proxy import reinitialize_upstreams +from routstr.upstream.model_paths import encode_model_path # The conftest patches ``routstr.payment.price.sats_usd_price``, but @@ -174,6 +175,33 @@ def _mock_chat_response( } +def _caching_upstream(*, cached_tokens: int = 2900, report_cost: bool = True) -> Any: + """Side effect that answers the one-token probe and then two long + prompts, reporting a cache hit (and optionally a USD cost) on the + repeated one — the shape an OpenAI-compatible caching upstream returns.""" + long_calls = {"n": 0} + + def _respond(request: Any) -> Response: + body = json.loads(request.content) + is_long = body["messages"][0]["role"] == "system" + if not is_long: + usage: dict[str, Any] = {"prompt_tokens": 5, "completion_tokens": 1} + if report_cost: + usage["cost"] = 9e-7 + else: + long_calls["n"] += 1 + usage = {"prompt_tokens": 3000, "completion_tokens": 1} + if long_calls["n"] > 1 and cached_tokens: + usage["prompt_tokens_details"] = {"cached_tokens": cached_tokens} + if report_cost: + usage["cost"] = 5e-5 if long_calls["n"] > 1 else 4e-4 + payload = _mock_chat_response() + payload["usage"] = usage + return Response(200, json=payload) + + return _respond + + @pytest.mark.integration @pytest.mark.asyncio async def test_certify_requires_admin_auth( @@ -205,13 +233,17 @@ async def test_certify_unknown_provider_returns_404( async def test_certify_all_ok( integration_client: AsyncClient, integration_session: AsyncSession ) -> None: - provider_id = await _seed_and_init(integration_session, integration_client) + provider_id = await _seed_and_init( + integration_session, + integration_client, + pricing_overrides={"input_cache_read": 1.4e-8}, + ) respx.get("https://certify-upstream.example/v1/models").mock( return_value=Response(200, json=_mock_models_response()) ) respx.post("https://certify-upstream.example/v1/chat/completions").mock( - return_value=Response(200, json=_mock_chat_response()) + side_effect=_caching_upstream() ) resp = await integration_client.post( @@ -230,6 +262,9 @@ async def test_certify_all_ok( "endpoint.models_payload", "usage.capture", "cost.prompt_completion", + "cache.reported", + "cache.billing", + "cost.margin", ] for row_id in live_row_ids: row = _find_row(body["rows"], row_id) @@ -239,6 +274,190 @@ async def test_certify_all_ok( assert item["status"] == "ok", f"{item['goal']}: {item}" +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_model_path_pins_every_completion( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + base_url = "https://openrouter.ai/api/v1" + provider_id = await _seed_and_init( + integration_session, + integration_client, + pricing_overrides={"input_cache_read": 1.4e-8}, + base_url=base_url, + ) + model_path = encode_model_path(base_url, "cert-test-model", "azure") + integration_session.add( + ModelPathRow( + model_id="cert-test-model", + path=model_path, + provider_slug="openrouter", + provider_type="openrouter", + endpoint_tag="azure", + endpoint_name="Azure", + model_metadata="{}", + upstream_provider_id=provider_id, + ) + ) + await integration_session.commit() + + respx.get(f"{base_url}/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + chat = respx.post(f"{base_url}/chat/completions").mock( + side_effect=_caching_upstream() + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={"model_id": "cert-test-model", "model_path": model_path}, + ) + assert resp.status_code == 200, resp.text + assert chat.call_count == 3 + for call in chat.calls: + body = json.loads(call.request.content) + assert body["provider"] == { + "order": ["azure"], + "allow_fallbacks": False, + } + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_uses_selected_path_pricing_for_margin( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + base_url = "https://openrouter.ai/api/v1" + provider_id = await _seed_and_init( + integration_session, + integration_client, + provider_fee=0.4, + pricing_overrides={ + "prompt": 1e-7, + "completion": 5e-7, + "input_cache_read": 1e-8, + }, + base_url=base_url, + ) + model_path = encode_model_path( + base_url, "cert-test-model", "deepinfra/fp8" + ) + integration_session.add( + ModelPathRow( + model_id="cert-test-model", + path=model_path, + provider_slug="openrouter", + provider_type="openrouter", + endpoint_tag="deepinfra/fp8", + endpoint_name="DeepInfra", + model_metadata=json.dumps( + { + "id": "cert-test-model", + "pricing": { + "prompt": 1.4e-7, + "completion": 4.2e-7, + "input_cache_read": 4.2e-9, + }, + } + ), + upstream_provider_id=provider_id, + ) + ) + await integration_session.commit() + + respx.get(f"{base_url}/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + usages = iter( + [ + { + "prompt_tokens": 31, + "completion_tokens": 1, + "cost": 0.00000459, + }, + { + "prompt_tokens": 4442, + "completion_tokens": 1, + "cost": 0.00057802, + }, + { + "prompt_tokens": 4442, + "completion_tokens": 1, + "prompt_tokens_details": {"cached_tokens": 4352}, + "cost": 0.00005578, + }, + ] + ) + + def _respond(_request: Any) -> Response: + payload = _mock_chat_response() + payload["usage"] = next(usages) + return Response(200, json=payload) + + respx.post(f"{base_url}/chat/completions").mock(side_effect=_respond) + sats_usd = 0.0008616302499999999 + with patch("routstr.payment.price.sats_usd_price", return_value=sats_usd): + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={"model_id": "cert-test-model", "model_path": model_path}, + ) + + assert resp.status_code == 200, resp.text + margin = _find_row(resp.json()["rows"], "cost.margin") + assert [ + (sample["upstream_msats_with_fee"], sample["configured_msats"]) + for sample in margin["evidence"]["samples"] + ] == [(3, 3), (269, 289), (26, 15)] + assert "289 < 269" not in margin["detail"] + assert "15 < 26" in margin["detail"] + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_provider_models_includes_certification_paths( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + base_url = "https://openrouter.ai/api/v1" + provider_id = await _seed_and_init( + integration_session, integration_client, base_url=base_url + ) + model_path = encode_model_path(base_url, "cert-test-model", "azure") + integration_session.add( + ModelPathRow( + model_id="cert-test-model", + path=model_path, + provider_slug="openrouter", + provider_type="openrouter", + endpoint_tag="azure", + endpoint_name="Azure", + model_metadata="{}", + upstream_provider_id=provider_id, + ) + ) + await integration_session.commit() + respx.get(f"{base_url}/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + + resp = await integration_client.get( + f"/admin/api/upstream-providers/{provider_id}/models", + headers=_admin_headers(), + ) + assert resp.status_code == 200, resp.text + assert resp.json()["certification_paths"]["cert-test-model"] == [ + { + "path": model_path, + "endpoint_tag": "azure", + "endpoint_name": "Azure", + } + ] + + @pytest.mark.integration @pytest.mark.asyncio @respx.mock @@ -527,6 +746,76 @@ async def test_certify_with_explicit_model_id( assert usage_row["status"] == "ok" +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_explicit_discovered_model_without_override( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + """An operator can probe a discovered model before creating an override.""" + from routstr.payment.models import ( + Architecture, + Model, + Pricing, + _update_model_sats_pricing, + ) + + provider = await _make_provider(integration_session) + assert provider.id is not None + remote_model = _update_model_sats_pricing( + Model( + id="remote-model", + name="Remote model", + description="", + created=0, + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=1e-7, completion=2e-7), + sats_pricing=None, + per_request_limits=None, + top_provider=None, + enabled=True, + upstream_provider_id=provider.id, + canonical_slug=None, + ), + 0.0005, + ) + + class FakeUpstream: + db_id = provider.id + + def get_cached_models(self) -> list[Model]: + return [remote_model] + + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response( + 200, json=_mock_models_response(models=[{"id": "remote-model"}]) + ) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response(model="remote-model")) + ) + + with patch("routstr.proxy.get_upstreams", return_value=[FakeUpstream()]): + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={"model_id": "remote-model", "check_cache": False}, + ) + + assert resp.status_code == 200, resp.text + rows = resp.json()["rows"] + assert _find_row(rows, "endpoint.reachable")["status"] == "ok" + assert _find_row(rows, "usage.capture")["status"] == "ok" + assert _find_row(rows, "cost.prompt_completion")["status"] == "ok" + + @pytest.mark.integration @pytest.mark.asyncio @respx.mock @@ -595,3 +884,119 @@ async def test_certify_row_contract_shape( assert set(item) >= {"goal", "label", "status", "tick", "rows"} assert item["status"] in {"ok", "warn", "fail"} assert item["tick"] in {"☑️", "⚠️", "❌"} + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_cache_warn_when_upstream_never_hits( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + side_effect=_caching_upstream(cached_tokens=0) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200, resp.text + rows = resp.json()["rows"] + assert _find_row(rows, "cache.reported")["status"] == "warn" + assert _find_row(rows, "cache.billing")["status"] == "warn" + goals = {item["goal"]: item["status"] for item in resp.json()["checklist"]} + assert goals["caching"] == "warn" + assert goals["margin"] == "ok" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_cache_billing_warns_without_cache_rate( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + side_effect=_caching_upstream(report_cost=False) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200, resp.text + rows = resp.json()["rows"] + assert _find_row(rows, "cache.reported")["status"] == "ok" + billing = _find_row(rows, "cache.billing") + assert billing["status"] == "warn" + assert billing["evidence"]["actual_total_msats"] == 841 + assert _find_row(rows, "cost.margin")["status"] == "warn" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_margin_fails_when_upstream_costs_more( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + + def _expensive(request: Any) -> Response: + payload = _mock_chat_response() + payload["usage"] = {"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3} + return Response(200, json=payload) + + respx.post("https://certify-upstream.example/v1/chat/completions").mock( + side_effect=_expensive + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + assert resp.status_code == 200, resp.text + margin = _find_row(resp.json()["rows"], "cost.margin") + assert margin["status"] == "fail" + goals = {item["goal"]: item["status"] for item in resp.json()["checklist"]} + assert goals["margin"] == "fail" + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_check_cache_false_skips_probe( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + respx.get("https://certify-upstream.example/v1/models").mock( + return_value=Response(200, json=_mock_models_response()) + ) + chat = respx.post("https://certify-upstream.example/v1/chat/completions").mock( + return_value=Response(200, json=_mock_chat_response()) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={"check_cache": False}, + ) + assert resp.status_code == 200, resp.text + assert chat.call_count == 1 + rows = resp.json()["rows"] + for row_id in ("cache.reported", "cache.billing", "cost.margin"): + row = _find_row(rows, row_id) + assert row["status"] == "warn" + assert "disabled" in row["detail"] diff --git a/tests/unit/test_certification.py b/tests/unit/test_certification.py index dc3ae9fc..1e963a79 100644 --- a/tests/unit/test_certification.py +++ b/tests/unit/test_certification.py @@ -527,9 +527,12 @@ class TestBuildChecklist: self._row("cost.prompt_completion", STATUS_OK), self._row("pricing.served_matches_configured", STATUS_OK), self._row("pricing.enabled_models_served", STATUS_OK), + self._row("cache.reported", STATUS_OK), + self._row("cache.billing", STATUS_OK), + self._row("cost.margin", STATUS_OK), ] checklist = build_checklist(rows) - assert len(checklist) == 4 + assert len(checklist) == 6 for item in checklist: assert item["status"] == STATUS_OK assert item["tick"] == "☑️" diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py new file mode 100644 index 00000000..2ca56482 --- /dev/null +++ b/tests/unit/test_certification_cache.py @@ -0,0 +1,499 @@ +"""Unit tests for the cache and margin rows in +routstr.upstream.certification_cache — no DB, network only via respx.""" + +from __future__ import annotations + +import json +from typing import Any + +import httpx +import pytest +import respx + +from routstr.payment.cost_calculation import CostData, CostDataError +from routstr.upstream.certification import STATUS_FAIL, STATUS_OK, STATUS_WARN +from routstr.upstream.certification_cache import ( + CacheProbeResult, + _raw_cache_keys, + cache_billing_row, + cache_probe_prefix, + cache_reported_row, + cost_margin_row, + probe_cache, + skipped_cache_rows, +) + +SATS_USD = 0.0005 +CHAT_URL = "https://upstream.example/v1/chat/completions" + + +def _model( + prompt: float = 1.4e-7, + completion: float = 2.8e-7, + cache_read: float = 0.0, + sats_usd: float = SATS_USD, +) -> Any: + from routstr.payment.models import ( + Architecture, + Model, + Pricing, + _update_model_sats_pricing, + ) + + model = Model( + id="test-model", + name="test-model", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing( + prompt=prompt, completion=completion, input_cache_read=cache_read + ), + sats_pricing=None, + per_request_limits=None, + top_provider=None, + enabled=True, + upstream_provider_id=None, + canonical_slug=None, + ) + return _update_model_sats_pricing(model, sats_usd) + + +def _payload(usage: dict[str, Any] | None) -> dict[str, Any]: + body: dict[str, Any] = {"choices": [{"message": {"content": "ok"}}]} + if usage is not None: + body["usage"] = usage + return body + + +def _probe( + payloads: list[dict[str, Any] | None], + statuses: list[int | None] | None = None, +) -> CacheProbeResult: + statuses = statuses if statuses is not None else [200] * len(payloads) + return CacheProbeResult( + chat_url=CHAT_URL, + statuses=statuses, + payloads=payloads, + errors=[None] * len(payloads), + latencies_ms=[1.0] * len(payloads), + ) + + +def _cost(total: int) -> CostData: + return CostData(base_msats=0, input_msats=total, output_msats=0, total_msats=total) + + +CACHED = { + "prompt_tokens": 3000, + "completion_tokens": 1, + "prompt_tokens_details": {"cached_tokens": 2900}, +} +UNCACHED = {"prompt_tokens": 3000, "completion_tokens": 1} + + +class TestPrefix: + def test_prefix_is_long_and_deterministic(self) -> None: + prefix = cache_probe_prefix() + assert len(prefix) > 8000 + assert prefix == cache_probe_prefix() + + +class TestRawCacheKeys: + def test_finds_nested_positive_cache_fields(self) -> None: + usage = {"prompt_tokens": 5, "details": {"cache_hits": 3, "cached": 0}} + assert _raw_cache_keys(usage) == ["details.cache_hits"] + + def test_ignores_bool_and_non_numeric(self) -> None: + assert _raw_cache_keys({"cached": True, "cache_key": "abc"}) == [] + + +class TestCacheReportedRow: + def test_ok_when_second_call_reports_cached_tokens(self) -> None: + row = cache_reported_row(_probe([_payload(UNCACHED), _payload(CACHED)])) + assert row["status"] == STATUS_OK + assert row["evidence"]["second_usage"]["cache_read_tokens"] == 2900 + + def test_warn_when_no_cache_hit(self) -> None: + row = cache_reported_row(_probe([_payload(UNCACHED), _payload(UNCACHED)])) + assert row["status"] == STATUS_WARN + + def test_fail_when_cache_reported_under_unknown_field(self) -> None: + second = _payload( + {"prompt_tokens": 3000, "completion_tokens": 1, "cache_hit_tokens": 2900} + ) + row = cache_reported_row(_probe([_payload(UNCACHED), second])) + assert row["status"] == STATUS_FAIL + assert row["evidence"]["unrecognised_cache_fields"] == ["cache_hit_tokens"] + + def test_fail_when_first_call_failed(self) -> None: + probe = CacheProbeResult( + chat_url=CHAT_URL, + statuses=[None], + payloads=[None], + errors=["ConnectError: boom"], + latencies_ms=[1.0], + ) + row = cache_reported_row(probe) + assert row["status"] == STATUS_FAIL + assert "ConnectError" in row["detail"] + + def test_fail_when_second_call_non_2xx(self) -> None: + row = cache_reported_row( + _probe([_payload(UNCACHED), {"error": "rate limited"}], [200, 429]) + ) + assert row["status"] == STATUS_FAIL + + +class TestCacheBillingRow: + def test_warn_when_nothing_cached(self) -> None: + row = cache_billing_row( + model=_model(), probe=_probe([_payload(UNCACHED)] * 2), cost_data=_cost(1) + ) + assert row["status"] == STATUS_WARN + + def test_warn_when_pricing_unknown(self) -> None: + row = cache_billing_row( + model=_model(), + probe=_probe([_payload(UNCACHED), _payload(CACHED)]), + cost_data=_cost(1), + pricing_known=False, + ) + assert row["status"] == STATUS_WARN + + def test_fail_on_engine_error(self) -> None: + row = cache_billing_row( + model=_model(), + probe=_probe([_payload(UNCACHED), _payload(CACHED)]), + cost_data=CostDataError(message="nope", code="pricing_error"), + ) + assert row["status"] == STATUS_FAIL + + def test_ok_with_discounted_rate(self) -> None: + # 280 msats/1k input, 28 msats/1k cached, 560 msats/1k output: + # 100*0.28 + 2900*0.028 + 1*0.56 = 109.76 -> 110 + row = cache_billing_row( + model=_model(cache_read=1.4e-8), + probe=_probe([_payload(UNCACHED), _payload(CACHED)]), + cost_data=_cost(110), + ) + assert row["status"] == STATUS_OK, row + assert row["evidence"]["full_price_total_msats"] == 841 + + def test_warn_when_cached_billed_at_full_rate(self) -> None: + row = cache_billing_row( + model=_model(), + probe=_probe([_payload(UNCACHED), _payload(CACHED)]), + cost_data=_cost(841), + ) + assert row["status"] == STATUS_WARN + assert "full input rate" in row["detail"] + + def test_fail_when_engine_disagrees(self) -> None: + row = cache_billing_row( + model=_model(cache_read=1.4e-8), + probe=_probe([_payload(UNCACHED), _payload(CACHED)]), + cost_data=_cost(841), + ) + assert row["status"] == STATUS_FAIL + + def test_ok_when_upstream_reports_cost(self) -> None: + cached = _payload({**CACHED, "cost": 5e-5}) + row = cache_billing_row( + model=_model(), + probe=_probe([_payload(UNCACHED), cached]), + cost_data=_cost(100), + ) + assert row["status"] == STATUS_OK + assert row["evidence"]["reported_usd"] == 5e-5 + + +class TestCostMarginRow: + def _row( + self, payloads: list[Any], fee: float = 1.0, model: Any = None + ) -> dict[str, Any]: + return cost_margin_row( + model=model or _model(), + payloads=payloads, + provider_fee=fee, + sats_to_usd=SATS_USD, + ) + + def test_warn_when_no_cost_reported(self) -> None: + row = self._row([_payload(UNCACHED), None]) + assert row["status"] == STATUS_WARN + assert row["evidence"]["samples"] == [] + + def test_ok_when_configured_covers_upstream(self) -> None: + # configured: 5*0.28 + 1*0.56 = 1.96 -> 2 msats; upstream 9e-7 USD -> 2 + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = self._row([payload]) + assert row["status"] == STATUS_OK, row + assert row["evidence"]["samples"][0]["upstream_msats_with_fee"] == 2 + + def test_fail_when_upstream_costs_more(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3}) + row = self._row([payload]) + assert row["status"] == STATUS_FAIL + assert "lose money" in row["detail"] + + def test_fee_scales_upstream_cost(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + assert self._row([payload], fee=1.0)["status"] == STATUS_OK + assert self._row([payload], fee=2.0)["status"] == STATUS_FAIL + + def test_deepseek_endpoint_price_exceeds_configured_model_price(self) -> None: + sats_usd = 0.0008616302499999999 + model = _model( + prompt=4e-8, + completion=2e-7, + cache_read=4e-9, + sats_usd=sats_usd, + ) + payloads = [ + _payload( + { + "prompt_tokens": 31, + "completion_tokens": 1, + "cost": 0.00000459, + } + ), + _payload( + { + "prompt_tokens": 4442, + "completion_tokens": 1, + "cost": 0.00057802, + } + ), + _payload( + { + "prompt_tokens": 4442, + "completion_tokens": 1, + "prompt_tokens_details": {"cached_tokens": 4352}, + "cost": 0.00005578, + } + ), + ] + + row = cost_margin_row( + model=model, + payloads=payloads, + provider_fee=0.4, + sats_to_usd=sats_usd, + ) + + assert row["status"] == STATUS_FAIL + assert [ + (sample["upstream_msats_with_fee"], sample["configured_msats"]) + for sample in row["evidence"]["samples"] + ] == [(3, 2), (269, 207), (26, 25)] + assert "207 < 269" in row["detail"] + assert "2 < 3" not in row["detail"] + assert "25 < 26" not in row["detail"] + + def test_one_msat_margin_gap_is_rounding_tolerance(self) -> None: + sats_usd = 0.0008616302499999999 + model = _model( + prompt=4e-8, + completion=2e-7, + cache_read=4e-9, + sats_usd=sats_usd, + ) + payloads = [ + _payload( + { + "prompt_tokens": 31, + "completion_tokens": 1, + "cost": 0.00000459, + } + ), + _payload( + { + "prompt_tokens": 4442, + "completion_tokens": 1, + "prompt_tokens_details": {"cached_tokens": 4352}, + "cost": 0.00005578, + } + ), + ] + + row = cost_margin_row( + model=model, + payloads=payloads, + provider_fee=0.4, + sats_to_usd=sats_usd, + ) + + assert row["status"] == STATUS_OK + + def test_warn_when_pricing_unknown(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + pricing_known=False, + ) + assert row["status"] == STATUS_WARN + + +class TestSkippedRows: + def test_three_warn_rows(self) -> None: + rows = skipped_cache_rows("Skipped") + assert [r["id"] for r in rows] == [ + "cache.reported", + "cache.billing", + "cost.margin", + ] + assert all(r["status"] == STATUS_WARN for r in rows) + + +class TestProbeCache: + @pytest.mark.asyncio + @respx.mock + async def test_sends_same_prompt_twice_with_cache_control(self) -> None: + route = respx.post(CHAT_URL).mock( + return_value=httpx.Response(200, json=_payload(CACHED)) + ) + async with httpx.AsyncClient() as client: + result = await probe_cache( + "https://upstream.example/v1", "k", "m", client=client + ) + assert result.request_format == "cache_control" + assert route.call_count == 2 + first, second = (json.loads(c.request.content) for c in route.calls) + assert first == second + system = first["messages"][0]["content"] + assert system[0]["cache_control"] == {"type": "ephemeral"} + assert first["max_tokens"] == 1 + assert route.calls[0].request.headers["Authorization"] == "Bearer k" + + @pytest.mark.asyncio + @respx.mock + async def test_pins_both_requests_to_one_endpoint(self) -> None: + route = respx.post(CHAT_URL).mock( + return_value=httpx.Response(200, json=_payload(CACHED)) + ) + async with httpx.AsyncClient() as client: + await probe_cache( + "https://upstream.example/v1", + "", + "m", + endpoint_tag="azure/swedencentral", + client=client, + ) + assert route.call_count == 2 + for call in route.calls: + body = json.loads(call.request.content) + assert body["provider"] == { + "order": ["azure/swedencentral"], + "allow_fallbacks": False, + } + + @pytest.mark.asyncio + @respx.mock + async def test_falls_back_to_plain_system_on_rejection(self) -> None: + route = respx.post(CHAT_URL).mock( + side_effect=[ + httpx.Response(400, json={"error": "cache_control not allowed"}), + httpx.Response(200, json=_payload(UNCACHED)), + httpx.Response(200, json=_payload(CACHED)), + ] + ) + async with httpx.AsyncClient() as client: + result = await probe_cache( + "https://upstream.example/v1", "", "m", client=client + ) + assert result.request_format == "plain" + assert route.call_count == 3 + assert result.statuses == [200, 200] + body = json.loads(route.calls[1].request.content) + assert isinstance(body["messages"][0]["content"], str) + + @pytest.mark.asyncio + @respx.mock + async def test_stops_after_failed_first_call(self) -> None: + route = respx.post(CHAT_URL).mock(side_effect=httpx.ConnectError("boom")) + async with httpx.AsyncClient() as client: + result = await probe_cache( + "https://upstream.example/v1", "", "m", client=client + ) + assert route.call_count == 1 + assert result.payloads == [None] + assert "ConnectError" in (result.second_error or "") + + +class TestErrorBranches: + @pytest.mark.asyncio + @respx.mock + async def test_non_json_and_list_bodies_are_errors(self) -> None: + respx.post(CHAT_URL).mock( + side_effect=[ + httpx.Response(200, content=b"not json"), + httpx.Response(200, json=[1, 2]), + ] + ) + async with httpx.AsyncClient() as client: + result = await probe_cache( + "https://upstream.example/v1", "", "m", client=client + ) + assert result.payloads == [None, None] + assert result.errors[0] is not None + assert "list" in (result.errors[1] or "") + + def test_malformed_usage_object_is_not_a_hit(self) -> None: + second = _payload({"prompt_tokens": {"nested": True}}) + row = cache_reported_row(_probe([_payload(UNCACHED), second])) + assert row["status"] == STATUS_WARN + + def test_billing_fails_on_non_finite_rate(self) -> None: + row = cache_billing_row( + model=_model(prompt=float("inf"), cache_read=1.4e-8), + probe=_probe([_payload(UNCACHED), _payload(CACHED)]), + cost_data=_cost(1), + ) + assert row["status"] == STATUS_FAIL + assert "error" in row["evidence"] + + def test_margin_fails_on_zero_sats_price(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = cost_margin_row( + model=_model(), payloads=[payload], provider_fee=1.0, sats_to_usd=0.0 + ) + assert row["status"] == STATUS_FAIL + + @pytest.mark.asyncio + async def test_engine_raise_becomes_billing_fail(self) -> None: + from unittest.mock import patch + + from routstr.upstream.certification_cache import run_cache_checks + + async def _boom(*args: Any, **kwargs: Any) -> Any: + raise RuntimeError("engine exploded") + + with ( + patch( + "routstr.upstream.certification_cache.probe_cache", + return_value=_probe([_payload(UNCACHED), _payload(CACHED)]), + ), + patch("routstr.upstream.certification_cache.calculate_cost", _boom), + ): + rows = await run_cache_checks( + "https://upstream.example/v1", + "", + _model(), + provider_fee=1.0, + sats_to_usd=SATS_USD, + probe_payload=_payload(UNCACHED), + ) + by_id = {r["id"]: r for r in rows} + assert by_id["cache.billing"]["status"] == STATUS_FAIL + assert "engine exploded" in by_id["cache.billing"]["detail"] diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py index 52f2fb02..7eb05552 100644 --- a/tests/unit/test_certification_hardening.py +++ b/tests/unit/test_certification_hardening.py @@ -337,8 +337,8 @@ class TestCliFreshProcess: ) 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 + assert len(document[0]["rows"]) == 8 + assert len(document[0]["checklist"]) == 6 def test_dead_host_exits_non_zero(self) -> None: result = _run_cli("--url", "http://localhost:1/v1", "--timeout", "1") diff --git a/ui/components/provider-card.tsx b/ui/components/provider-card.tsx index 44f05183..ef07013a 100644 --- a/ui/components/provider-card.tsx +++ b/ui/components/provider-card.tsx @@ -25,9 +25,11 @@ import { AlertTriangle, Unlock, Loader2, + ShieldCheck, } from 'lucide-react'; import { ProviderBalance } from '@/components/provider-balance'; import { ProviderModelsPanel } from '@/components/provider-models-panel'; +import { ProviderCertificationDialog } from '@/components/provider-certification-dialog'; import { RoutstrCreateKeySection } from '@/components/providers/RoutstrCreateKeySection'; import { RoutstrProviderService } from '@/lib/api/services/routstr-provider'; import { getErrorStatus } from '@/lib/api/client'; @@ -93,6 +95,7 @@ export function ProviderCard({ }: ProviderCardProps) { const queryClient = useQueryClient(); const [isKeyModalOpen, setIsKeyModalOpen] = useState(false); + const [isCertifyOpen, setIsCertifyOpen] = useState(false); const [isReleaseDialogOpen, setIsReleaseDialogOpen] = useState(false); // The claim as the query cache held it when the admin opened the dialog. // The mutation sends this token rather than re-reading the query at submit @@ -310,6 +313,17 @@ export function ProviderCard({ )} + + + {showEvidence && ( +
+              {JSON.stringify(row.evidence, null, 2)}
+            
+ )} + + )} + + ); +} + +function ChecklistSummary({ report }: { report: ProviderCertification }) { + return ( +
    + {report.checklist.map((goal) => { + const style = STATUS_STYLES[goal.status]; + const Icon = style.icon; + return ( +
  • + + {goal.label} +
  • + ); + })} +
+ ); +} + +function CertificationReport({ report }: { report: ProviderCertification }) { + const failing = report.rows.filter((row) => row.status === 'fail').length; + const warning = report.rows.filter((row) => row.status === 'warn').length; + + return ( +
+ +
+ + {report.rows.length} checks · {failing} failed · {warning} warnings + + Generated {new Date(report.generated_at).toLocaleString()} +
+
    + {report.rows.map((row) => ( + + ))} +
+
+ ); +} + +function getErrorMessage(error: unknown): string { + if (error instanceof Error) return error.message; + return 'Certification request failed'; +} + +function resultStatus( + result: ModelCertificationResult +): CertificationStatus | 'error' { + if (result.error || !result.report) return 'error'; + if (result.report.rows.some((row) => row.status === 'fail')) return 'fail'; + if (result.report.rows.some((row) => row.status === 'warn')) return 'warn'; + return 'ok'; +} + +export function ProviderCertificationDialog({ + provider, + open, + onOpenChange, +}: ProviderCertificationDialogProps) { + const [checkCache, setCheckCache] = useState(true); + const [selectedModelIds, setSelectedModelIds] = useState([]); + const [pathModes, setPathModes] = useState>({}); + const [selectedModelPaths, setSelectedModelPaths] = useState< + Record + >({}); + const [results, setResults] = useState([]); + const [currentModel, setCurrentModel] = useState<{ + id: string; + index: number; + total: number; + pathCount: number; + } | null>(null); + + const models = useQuery({ + queryKey: ['provider-models', provider.id], + queryFn: () => AdminService.getProviderModels(provider.id), + enabled: open, + }); + + const certify = useMutation({ + mutationFn: async ({ + modelRuns, + includeCache, + }: { + modelRuns: ModelRun[]; + includeCache: boolean; + }) => { + const completed: ModelCertificationResult[] = []; + setResults([]); + + for (const [index, run] of modelRuns.entries()) { + setCurrentModel({ + id: run.modelId, + index: index + 1, + total: modelRuns.length, + pathCount: run.targets.length, + }); + const batch = await Promise.all( + run.targets.map(async (target): Promise => { + const resultKey = `${run.modelId}::${target.path ?? 'default'}`; + try { + const report = await AdminService.certifyProvider(provider.id, { + model_id: run.modelId, + model_path: target.path, + check_cache: includeCache, + }); + return { + resultKey, + modelId: run.modelId, + pathLabel: target.label, + report, + }; + } catch (error) { + return { + resultKey, + modelId: run.modelId, + pathLabel: target.label, + error: getErrorMessage(error), + }; + } + }) + ); + completed.push(...batch); + setResults([...completed]); + } + + return completed; + }, + onSettled: () => setCurrentModel(null), + }); + const { reset: resetCertification } = certify; + + useEffect(() => { + if (!open) { + resetCertification(); + setSelectedModelIds([]); + setPathModes({}); + setSelectedModelPaths({}); + setResults([]); + setCurrentModel(null); + } + }, [open, resetCertification]); + + const configuredOptions: ModelOption[] = + models.data?.db_models.map((model) => ({ + model, + source: 'configured', + })) ?? []; + const discoveredOptions: ModelOption[] = + models.data?.remote_models.map((model) => ({ + model, + source: 'discovered', + })) ?? []; + const allOptions = [...configuredOptions, ...discoveredOptions]; + const namesById = new Map( + allOptions.map(({ model }) => [model.id, model.name || model.id]) + ); + + const pathsForModel = (modelId: string): CertificationPath[] => + (models.data?.certification_paths[modelId] ?? []).filter( + (path) => path.endpoint_tag + ); + + const pathLabel = (path: CertificationPath): string => + path.endpoint_name && path.endpoint_name !== path.endpoint_tag + ? `${path.endpoint_name} (${path.endpoint_tag})` + : path.endpoint_tag || 'Provider default'; + + const toggleModel = (modelId: string) => { + const isSelected = selectedModelIds.includes(modelId); + setSelectedModelIds((current) => + isSelected + ? current.filter((id) => id !== modelId) + : [...current, modelId] + ); + setPathModes((current) => { + const next = { ...current }; + if (isSelected) delete next[modelId]; + else next[modelId] = 'default'; + return next; + }); + setSelectedModelPaths((current) => { + const next = { ...current }; + if (isSelected) delete next[modelId]; + return next; + }); + resetCertification(); + setResults([]); + }; + + const modelsNeedingPath = selectedModelIds.filter( + (modelId) => + pathModes[modelId] === 'selected' && + (selectedModelPaths[modelId]?.length ?? 0) === 0 + ); + + const buildModelRuns = (): ModelRun[] => + selectedModelIds.map((modelId) => { + const paths = pathsForModel(modelId); + const mode = pathModes[modelId] ?? 'default'; + if (mode === 'all') { + return { + modelId, + targets: paths.map((path) => ({ + path: path.path, + label: pathLabel(path), + })), + }; + } + if (mode === 'selected') { + const selected = new Set(selectedModelPaths[modelId] ?? []); + return { + modelId, + targets: paths + .filter((path) => selected.has(path.path)) + .map((path) => ({ path: path.path, label: pathLabel(path) })), + }; + } + return { modelId, targets: [{ label: 'Provider default' }] }; + }); + + const modelRuns = buildModelRuns(); + const targetCount = modelRuns.reduce( + (total, run) => total + run.targets.length, + 0 + ); + + const renderModelGroup = (label: string, options: ModelOption[]) => { + if (options.length === 0) return null; + return ( + + {options.map(({ model }) => ( + toggleModel(model.id)} + disabled={certify.isPending} + > + + ))} + + ); + }; + + return ( + + + + Certify upstream models + + Select models, then use the provider default, choose specific + paths, or test every path. Models run one at a time; paths for the + same model run in parallel against {provider.base_url}. + + + +
+
+ +
+ + {selectedModelIds.length} selected + + {selectedModelIds.length > 0 && !certify.isPending && ( + + )} +
+
+ + + + + {models.isLoading ? 'Loading models…' : 'No models found'} + + {renderModelGroup('Configured models', configuredOptions)} + {renderModelGroup('Discovered models', discoveredOptions)} + + + {models.isError && ( +

+ {getErrorMessage(models.error)} +

+ )} + {selectedModelIds.map((modelId) => { + const paths = pathsForModel(modelId); + const mode = pathModes[modelId] ?? 'default'; + const selectedPaths = selectedModelPaths[modelId] ?? []; + return ( +
+
+
+ {namesById.get(modelId) ?? modelId} +
+
+ {modelId} +
+
+ {paths.length > 0 ? ( + <> + { + if (!value) return; + setPathModes((current) => ({ + ...current, + [modelId]: value as ModelPathMode, + })); + }} + disabled={certify.isPending} + className='w-full justify-start' + > + Default + Choose paths + All paths + + {mode === 'default' && ( +

+ Uses the upstream provider's normal model routing. +

+ )} + {mode === 'selected' && ( + + + + + +
+ {paths.map((path) => { + const checked = selectedPaths.includes(path.path); + return ( + + ); + })} +
+
+
+ )} + {mode === 'all' && ( +

+ All {paths.length} paths will run in parallel. +

+ )} + + ) : ( +

+ Only the provider default route is available. +

+ )} +
+ ); + })} + {modelsNeedingPath.length > 0 && ( +

+ Choose at least one path for each model using “Choose paths”. +

+ )} +
+ +
+
+ setCheckCache(value === true)} + disabled={certify.isPending} + /> + +
+ +
+ + {currentModel && ( +
+ + Probing {namesById.get(currentModel.id) ?? currentModel.id} + {currentModel.pathCount > 1 + ? ` across ${currentModel.pathCount} paths in parallel` + : ''}{' '} + — model {currentModel.index} of {currentModel.total} +
+ )} + + {results.length > 0 && ( + result.resultKey).join('|')} + defaultValue={results[0].resultKey} + className='space-y-3' + > +
+ + {results.map((result) => { + const status = resultStatus(result); + const Icon = + status === 'error' ? XCircle : STATUS_STYLES[status].icon; + return ( + + + + {namesById.get(result.modelId) ?? result.modelId} ·{' '} + {result.pathLabel} + + + ); + })} + +
+ {results.map((result) => ( + + {result.report ? ( + + ) : ( +
+ {result.error ?? 'Certification failed'} +
+ )} +
+ ))} +
+ )} +
+
+ ); +} diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index f21e7b64..ba665da0 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -46,6 +46,43 @@ export const UpdateUpstreamProviderSchema = z.object({ slug: z.string().optional(), }); +export const CertificationStatusSchema = z.enum(['ok', 'warn', 'fail']); + +export const CertificationRowSchema = z.object({ + id: z.string(), + status: CertificationStatusSchema, + title: z.string(), + detail: z.string(), + evidence: z.record(z.string(), z.unknown()), +}); + +export const CertificationGoalSchema = z.object({ + goal: z.string(), + label: z.string(), + status: CertificationStatusSchema, + tick: z.string(), + rows: z.array(z.string()), +}); + +export const ProviderCertificationSchema = z.object({ + provider_id: z.number(), + generated_at: z.string(), + rows: z.array(CertificationRowSchema), + checklist: z.array(CertificationGoalSchema), +}); + +export type CertificationStatus = z.infer; +export type CertificationRow = z.infer; +export type CertificationGoal = z.infer; +export type ProviderCertification = z.infer; + +export type CertifyProviderRequest = { + model_id?: string; + model_path?: string; + timeout_seconds?: number; + check_cache?: boolean; +}; + export const AdminModelPricingSchema = z.object({ prompt: z.number().optional(), completion: z.number().optional(), @@ -84,6 +121,12 @@ export const AdminModelSchema = z.object({ forwarded_model_id: z.string().nullable().optional(), }); +export const CertificationPathSchema = z.object({ + path: z.string(), + endpoint_tag: z.string().nullable(), + endpoint_name: z.string().nullable(), +}); + export const ProviderModelsSchema = z.object({ provider: z.object({ id: z.number(), @@ -92,6 +135,10 @@ export const ProviderModelsSchema = z.object({ }), db_models: z.array(AdminModelSchema), remote_models: z.array(AdminModelSchema), + certification_paths: z.record( + z.string(), + z.array(CertificationPathSchema) + ), }); export type ProviderType = z.infer; @@ -107,6 +154,7 @@ export type AdminModelPricing = z.infer; export type AdminModelArchitecture = z.infer< typeof AdminModelArchitectureSchema >; +export type CertificationPath = z.infer; export type ProviderModels = z.infer; export interface AdminModelAsModel { @@ -317,6 +365,17 @@ export class AdminService { ); } + static async certifyProvider( + providerId: number, + body: CertifyProviderRequest = {} + ): Promise { + const data = await apiClient.post( + `/admin/api/upstream-providers/${providerId}/certify`, + body + ); + return ProviderCertificationSchema.parse(data); + } + static async getProviderModels(providerId: number): Promise { const data = await apiClient.get( `/admin/api/upstream-providers/${providerId}/models` From 96eef3c6b4c6ad3449529bec9767c5a70fa41a2b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 21:31:39 +0200 Subject: [PATCH 09/18] feat: improve certification workspace --- .../provider-certification-dialog.tsx | 609 ++++++++++-------- 1 file changed, 355 insertions(+), 254 deletions(-) diff --git a/ui/components/provider-certification-dialog.tsx b/ui/components/provider-certification-dialog.tsx index bfe988bf..afa34d29 100644 --- a/ui/components/provider-certification-dialog.tsx +++ b/ui/components/provider-certification-dialog.tsx @@ -221,6 +221,7 @@ export function ProviderCertificationDialog({ onOpenChange, }: ProviderCertificationDialogProps) { const [checkCache, setCheckCache] = useState(true); + const [workspaceTab, setWorkspaceTab] = useState<'setup' | 'results'>('setup'); const [selectedModelIds, setSelectedModelIds] = useState([]); const [pathModes, setPathModes] = useState>({}); const [selectedModelPaths, setSelectedModelPaths] = useState< @@ -250,6 +251,7 @@ export function ProviderCertificationDialog({ }) => { const completed: ModelCertificationResult[] = []; setResults([]); + setWorkspaceTab('results'); for (const [index, run] of modelRuns.entries()) { setCurrentModel({ @@ -296,6 +298,7 @@ export function ProviderCertificationDialog({ useEffect(() => { if (!open) { resetCertification(); + setWorkspaceTab('setup'); setSelectedModelIds([]); setPathModes({}); setSelectedModelPaths({}); @@ -419,8 +422,8 @@ export function ProviderCertificationDialog({ return ( - - + + Certify upstream models Select models, then use the provider default, choose specific @@ -429,268 +432,366 @@ export function ProviderCertificationDialog({ -
-
- -
- - {selectedModelIds.length} selected - - {selectedModelIds.length > 0 && !certify.isPending && ( - + + setWorkspaceTab(value as 'setup' | 'results') + } + className='min-h-0 flex-1 overflow-hidden' + > + + Setup + + Results + {(certify.isPending || results.length > 0) && ( + + {results.length}/{targetCount} + )} -
-
- - - - - {models.isLoading ? 'Loading models…' : 'No models found'} - - {renderModelGroup('Configured models', configuredOptions)} - {renderModelGroup('Discovered models', discoveredOptions)} - - - {models.isError && ( -

- {getErrorMessage(models.error)} -

- )} - {selectedModelIds.map((modelId) => { - const paths = pathsForModel(modelId); - const mode = pathModes[modelId] ?? 'default'; - const selectedPaths = selectedModelPaths[modelId] ?? []; - return ( -
-
-
- {namesById.get(modelId) ?? modelId} -
-
- {modelId} + + + + +
+
+
+ +
+ + {selectedModelIds.length} selected + + {selectedModelIds.length > 0 && !certify.isPending && ( + + )}
- {paths.length > 0 ? ( - <> - { - if (!value) return; - setPathModes((current) => ({ - ...current, - [modelId]: value as ModelPathMode, - })); - }} - disabled={certify.isPending} - className='w-full justify-start' - > - Default - Choose paths - All paths - - {mode === 'default' && ( -

- Uses the upstream provider's normal model routing. -

- )} - {mode === 'selected' && ( - - - - - -
- {paths.map((path) => { - const checked = selectedPaths.includes(path.path); - return ( - - ); - })} -
-
-
- )} - {mode === 'all' && ( -

- All {paths.length} paths will run in parallel. -

- )} - - ) : ( -

- Only the provider default route is available. + + + + + {models.isLoading ? 'Loading models…' : 'No models found'} + + {renderModelGroup('Configured models', configuredOptions)} + {renderModelGroup('Discovered models', discoveredOptions)} + + + {models.isError && ( +

+ {getErrorMessage(models.error)}

)}
- ); - })} - {modelsNeedingPath.length > 0 && ( -

- Choose at least one path for each model using “Choose paths”. -

- )} -
-
-
- setCheckCache(value === true)} - disabled={certify.isPending} - /> - -
- -
- - {currentModel && ( -
- - Probing {namesById.get(currentModel.id) ?? currentModel.id} - {currentModel.pathCount > 1 - ? ` across ${currentModel.pathCount} paths in parallel` - : ''}{' '} - — model {currentModel.index} of {currentModel.total} -
- )} - - {results.length > 0 && ( - result.resultKey).join('|')} - defaultValue={results[0].resultKey} - className='space-y-3' - > -
- - {results.map((result) => { - const status = resultStatus(result); - const Icon = - status === 'error' ? XCircle : STATUS_STYLES[status].icon; - return ( - - { + const paths = pathsForModel(modelId); + const mode = pathModes[modelId] ?? 'default'; + const selectedPaths = selectedModelPaths[modelId] ?? []; + return ( +
+
+
+ {namesById.get(modelId) ?? modelId} +
+
+ {modelId} +
+
+ {paths.length > 0 ? ( + <> + { + if (!value) return; + setPathModes((current) => ({ + ...current, + [modelId]: value as ModelPathMode, + })); + }} + disabled={certify.isPending} + className='w-full justify-start' + > + Default + + Choose paths + + All paths + + {mode === 'default' && ( +

+ Uses the upstream provider's normal model routing. +

)} - /> - - {namesById.get(result.modelId) ?? result.modelId} ·{' '} - {result.pathLabel} - - - ); - })} - -
- {results.map((result) => ( - - {result.report ? ( - - ) : ( -
- {result.error ?? 'Certification failed'} + {mode === 'selected' && ( + + + + + +
+ {paths.map((path) => { + const checked = selectedPaths.includes( + path.path + ); + return ( + + ); + })} +
+
+
+ )} + {mode === 'all' && ( +

+ All {paths.length} paths will run in parallel. +

+ )} + + ) : ( +

+ Only the provider default route is available. +

+ )}
+ ); + })} + {modelsNeedingPath.length > 0 && ( +

+ Choose at least one path for each model using “Choose paths”. +

+ )} +
+ +
+
+ setCheckCache(value === true)} + disabled={certify.isPending} + /> + +
+ +
+
+ + + {currentModel && ( +
+ + + Probing {namesById.get(currentModel.id) ?? currentModel.id} + {currentModel.pathCount > 1 + ? ` across ${currentModel.pathCount} paths in parallel` + : ''}{' '} + — model {currentModel.index} of {currentModel.total} + +
+ )} + + {results.length === 0 ? ( +
+ Results will appear here as certification completes. +
+ ) : ( + result.resultKey).join('|')} + defaultValue={results[0].resultKey} + className='min-h-0 flex-1 overflow-hidden' + > +
+ + {results.map((result, index) => { + const status = resultStatus(result); + const Icon = + status === 'error' + ? XCircle + : STATUS_STYLES[status].icon; + const modelRouteNumber = results + .slice(0, index + 1) + .filter((item) => item.modelId === result.modelId).length; + const modelRouteCount = results.filter( + (item) => item.modelId === result.modelId + ).length; + return ( + + + + {namesById.get(result.modelId) ?? result.modelId} + + {modelRouteCount > 1 && ( + + {modelRouteNumber} + + )} + + ); + })} + +
+ +
+ {results.map((result) => { + const status = resultStatus(result); + return ( + +
+
+
+ {namesById.get(result.modelId) ?? result.modelId} +
+ {status === 'error' ? ( + + Error + + ) : ( + + )} +
+
+ + Model path + +
+ {result.pathLabel} +
+
+
+ {result.report ? ( + + ) : ( +
+ {result.error ?? 'Certification failed'} +
+ )} +
+ ); + })} +
+
+ )} +
+
); From 98bb904525d88751a2e0edb74cc66613404485ea Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 22:10:33 +0200 Subject: [PATCH 10/18] feat: add multi-provider certification page --- ui/app/providers/certification/page.tsx | 568 +++++++++++++ ui/app/providers/page.tsx | 21 +- .../provider-certification-dialog.tsx | 756 ++---------------- .../provider-certification-results.tsx | 305 +++++++ .../provider-certification-setup.tsx | 309 +++++++ ui/hooks/use-provider-certification-runner.ts | 128 +++ ui/lib/provider-certification.ts | 116 +++ 7 files changed, 1506 insertions(+), 697 deletions(-) create mode 100644 ui/app/providers/certification/page.tsx create mode 100644 ui/components/provider-certification-results.tsx create mode 100644 ui/components/provider-certification-setup.tsx create mode 100644 ui/hooks/use-provider-certification-runner.ts create mode 100644 ui/lib/provider-certification.ts diff --git a/ui/app/providers/certification/page.tsx b/ui/app/providers/certification/page.tsx new file mode 100644 index 00000000..ef1b0ef4 --- /dev/null +++ b/ui/app/providers/certification/page.tsx @@ -0,0 +1,568 @@ +'use client'; + +import { useEffect, useMemo, useState } from 'react'; +import Link from 'next/link'; +import { useQueries, useQuery } from '@tanstack/react-query'; +import { + ArrowLeft, + CheckCircle2, + ChevronDown, + Loader2, + Play, + RotateCcw, + Server, +} from 'lucide-react'; + +import { AppPageShell } from '@/components/app-page-shell'; +import { PageHeader } from '@/components/page-header'; +import { + ProviderCertificationResults, + summarizeCertificationResults, +} from '@/components/provider-certification-results'; +import { + getCertificationModelNames, + ProviderCertificationSetupPanel, +} from '@/components/provider-certification-setup'; +import { Badge } from '@/components/ui/badge'; +import { Button } from '@/components/ui/button'; +import { Card, CardContent } from '@/components/ui/card'; +import { Checkbox } from '@/components/ui/checkbox'; +import { + Command, + CommandEmpty, + CommandInput, + CommandItem, + CommandList, +} from '@/components/ui/command'; +import { + Popover, + PopoverContent, + PopoverTrigger, +} from '@/components/ui/popover'; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from '@/components/ui/select'; +import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; +import { runProviderCertification } from '@/hooks/use-provider-certification-runner'; +import { AdminService } from '@/lib/api/services/admin'; +import type { + ProviderModels, + UpstreamProvider, +} from '@/lib/api/services/admin'; +import { + buildModelRuns, + countCertificationTargets, + emptyCertificationSetup, + getModelsNeedingPath, +} from '@/lib/provider-certification'; +import type { + CertificationProgress, + ModelCertificationResult, + ProviderCertificationSetup, +} from '@/lib/provider-certification'; +import { cn } from '@/lib/utils'; + +function providerName(provider: UpstreamProvider): string { + return provider.slug || provider.provider_type; +} + +export default function MultiProviderCertificationPage() { + const [selectedProviderIds, setSelectedProviderIds] = useState([]); + const [activeProviderId, setActiveProviderId] = useState(null); + const [setups, setSetups] = useState< + Record + >({}); + const [workspaceTab, setWorkspaceTab] = useState<'setup' | 'results'>( + 'setup' + ); + const [resultsByProvider, setResultsByProvider] = useState< + Record + >({}); + const [progressByProvider, setProgressByProvider] = useState< + Record + >({}); + const [runningProviderIds, setRunningProviderIds] = useState([]); + + const providersQuery = useQuery({ + queryKey: ['upstream-providers'], + queryFn: () => AdminService.getUpstreamProviders(), + refetchOnWindowFocus: false, + }); + const providers = useMemo( + () => providersQuery.data ?? [], + [providersQuery.data] + ); + const providersById = useMemo( + () => new Map(providers.map((provider) => [provider.id, provider])), + [providers] + ); + + const modelQueries = useQueries({ + queries: selectedProviderIds.map((providerId) => ({ + queryKey: ['provider-models', providerId], + queryFn: () => AdminService.getProviderModels(providerId), + refetchOnWindowFocus: false, + })), + }); + const modelQueryByProvider = new Map( + selectedProviderIds.map((providerId, index) => [ + providerId, + modelQueries[index], + ]) + ); + + const selectedProviders = selectedProviderIds + .map((providerId) => providersById.get(providerId)) + .filter((provider): provider is UpstreamProvider => Boolean(provider)); + const activeProvider = activeProviderId + ? providersById.get(activeProviderId) + : undefined; + const activeSetup = activeProviderId + ? (setups[activeProviderId] ?? emptyCertificationSetup()) + : undefined; + const activeModelsQuery = activeProviderId + ? modelQueryByProvider.get(activeProviderId) + : undefined; + + useEffect(() => { + if (runningProviderIds.length === 0) return; + const warnBeforeUnload = (event: BeforeUnloadEvent) => { + event.preventDefault(); + }; + window.addEventListener('beforeunload', warnBeforeUnload); + return () => window.removeEventListener('beforeunload', warnBeforeUnload); + }, [runningProviderIds.length]); + + const toggleProvider = (providerId: number) => { + if (runningProviderIds.includes(providerId)) return; + const selected = selectedProviderIds.includes(providerId); + if (selected) { + const remaining = selectedProviderIds.filter((id) => id !== providerId); + setSelectedProviderIds(remaining); + if (activeProviderId === providerId) { + setActiveProviderId(remaining[0] ?? null); + } + return; + } + setSelectedProviderIds((current) => [...current, providerId]); + setSetups((current) => ({ + ...current, + [providerId]: current[providerId] ?? emptyCertificationSetup(), + })); + setActiveProviderId(providerId); + }; + + const updateProviderSetup = ( + providerId: number, + setup: ProviderCertificationSetup + ) => { + setSetups((current) => ({ ...current, [providerId]: setup })); + setResultsByProvider((current) => ({ ...current, [providerId]: [] })); + }; + + const providerRuns = (providerId: number) => + buildModelRuns( + setups[providerId] ?? emptyCertificationSetup(), + modelQueryByProvider.get(providerId)?.data as ProviderModels | undefined + ); + + const isProviderReady = (providerId: number): boolean => { + const setup = setups[providerId] ?? emptyCertificationSetup(); + return ( + setup.selectedModelIds.length > 0 && + getModelsNeedingPath(setup).length === 0 && + Boolean(modelQueryByProvider.get(providerId)?.data) + ); + }; + + const runOneProvider = async (providerId: number) => { + const setup = setups[providerId] ?? emptyCertificationSetup(); + const modelRuns = providerRuns(providerId); + setResultsByProvider((current) => ({ ...current, [providerId]: [] })); + setProgressByProvider((current) => ({ ...current, [providerId]: null })); + setRunningProviderIds((current) => + current.includes(providerId) ? current : [...current, providerId] + ); + try { + await runProviderCertification({ + providerId, + modelRuns, + includeCache: setup.checkCache, + onProgress: (progress) => + setProgressByProvider((current) => ({ + ...current, + [providerId]: progress, + })), + onResults: (results) => + setResultsByProvider((current) => ({ + ...current, + [providerId]: results, + })), + }); + } finally { + setProgressByProvider((current) => ({ + ...current, + [providerId]: null, + })); + setRunningProviderIds((current) => + current.filter((id) => id !== providerId) + ); + } + }; + + const runAllProviders = () => { + setWorkspaceTab('results'); + void Promise.all(selectedProviderIds.map(runOneProvider)); + }; + + const activeResults = activeProviderId + ? (resultsByProvider[activeProviderId] ?? []) + : []; + const activeNames = getCertificationModelNames( + activeModelsQuery?.data as ProviderModels | undefined + ); + + const totalModels = selectedProviderIds.reduce( + (total, providerId) => + total + (setups[providerId]?.selectedModelIds.length ?? 0), + 0 + ); + const totalRoutes = selectedProviderIds.reduce( + (total, providerId) => + total + countCertificationTargets(providerRuns(providerId)), + 0 + ); + const incompleteProviders = selectedProviderIds.filter( + (providerId) => !isProviderReady(providerId) + ); + const allReady = + selectedProviderIds.length > 0 && incompleteProviders.length === 0; + const allResults = Object.values(resultsByProvider).flat(); + const aggregateSummary = summarizeCertificationResults(allResults); + const pendingRoutes = Math.max(totalRoutes - allResults.length, 0); + + return ( + +
+ + + + + + + + + + + + {providersQuery.isLoading + ? 'Loading providers…' + : 'No providers found'} + + {providers.map((provider) => ( + toggleProvider(provider.id)} + disabled={runningProviderIds.includes(provider.id)} + > + + ))} + + + + +
+ } + /> + + {selectedProviders.length === 0 ? ( + + + +
+
Select providers to certify
+

+ Choose two or more providers to configure independent model + and path runs. +

+
+
+
+ ) : ( +
+ + +
+ +
+ + {activeProvider && activeSetup ? ( + +
+
+
+
+ {providerName(activeProvider)} +
+
+ {activeProvider.base_url} +
+
+ + {activeProvider.enabled ? 'Enabled' : 'Disabled'} + +
+
+ + + setWorkspaceTab(value as 'setup' | 'results') + } + className='min-h-0 flex-1 overflow-hidden px-4 pb-4' + > + + Setup + + Results + {activeResults.length > 0 && ( + + {activeResults.length} + + )} + + + + +
+ + updateProviderSetup(activeProvider.id, next) + } + disabled={runningProviderIds.includes( + activeProvider.id + )} + idPrefix={`multi-certify-${activeProvider.id}`} + /> +
+
+ +
+
+ + +
+ +
+ +
+
+
+ ) : null} +
+ )} + + {selectedProviders.length > 0 && ( +
+
+
+ {selectedProviderIds.length} providers · {totalModels} models ·{' '} + {totalRoutes} routes +
+
+ {incompleteProviders.length > 0 + ? `${incompleteProviders.length} provider${incompleteProviders.length === 1 ? '' : 's'} need a model or path selection. ` + : 'Providers run concurrently; models are sequential and paths run in parallel. '} + Status: {aggregateSummary.ok} ok, {aggregateSummary.warn}{' '} + warnings, {aggregateSummary.fail} failed,{' '} + {aggregateSummary.error} errors, {runningProviderIds.length}{' '} + running, {pendingRoutes} pending. +
+
+ +
+ )} + +
+ ); +} diff --git a/ui/app/providers/page.tsx b/ui/app/providers/page.tsx index 669ac088..250fcafa 100644 --- a/ui/app/providers/page.tsx +++ b/ui/app/providers/page.tsx @@ -1,5 +1,6 @@ 'use client'; +import Link from 'next/link'; import { Button } from '@/components/ui/button'; import { Card, CardContent } from '@/components/ui/card'; import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query'; @@ -18,7 +19,7 @@ import { BatchOverrideDialog } from '@/components/batch-override-dialog'; import { ProviderCard } from '@/components/provider-card'; import { ProviderFormDialogContent } from '@/components/provider-form-dialog-content'; import { Skeleton } from '@/components/ui/skeleton'; -import { AlertCircle, Plus, Server } from 'lucide-react'; +import { AlertCircle, BadgeCheck, Plus, Server } from 'lucide-react'; import { Alert, AlertDescription } from '@/components/ui/alert'; import { Dialog, DialogTrigger } from '@/components/ui/dialog'; import { @@ -348,12 +349,20 @@ export default function ProvidersPage() { title='Upstream Providers' description='Manage your AI provider connections and credentials.' actions={ - - - + + + + } /> void; } -interface ModelOption { - model: AdminModel; - source: 'configured' | 'discovered'; -} - -type ModelPathMode = 'default' | 'selected' | 'all'; - -interface ModelPathTarget { - path?: string; - label: string; -} - -interface ModelRun { - modelId: string; - targets: ModelPathTarget[]; -} - -interface ModelCertificationResult { - resultKey: string; - modelId: string; - pathLabel: string; - report?: ProviderCertification; - error?: string; -} - -const STATUS_STYLES: Record< - CertificationStatus, - { label: string; icon: typeof CheckCircle2; className: string } -> = { - ok: { - label: 'OK', - icon: CheckCircle2, - className: - 'border-emerald-500/40 bg-emerald-500/10 text-emerald-700 dark:text-emerald-400', - }, - warn: { - label: 'Warn', - icon: AlertTriangle, - className: - 'border-amber-500/40 bg-amber-500/10 text-amber-700 dark:text-amber-400', - }, - fail: { - label: 'Fail', - icon: XCircle, - className: 'border-red-500/40 bg-red-500/10 text-red-700 dark:text-red-400', - }, -}; - -function StatusBadge({ status }: { status: CertificationStatus }) { - const style = STATUS_STYLES[status]; - const Icon = style.icon; - return ( - - - {style.label} - - ); -} - -function RowItem({ row }: { row: CertificationRow }) { - const [showEvidence, setShowEvidence] = useState(false); - const hasEvidence = Object.keys(row.evidence).length > 0; - return ( -
  • -
    -
    -
    {row.title}
    -
    - {row.detail} -
    -
    - {row.id} -
    -
    - -
    - {hasEvidence && ( -
    - - {showEvidence && ( -
    -              {JSON.stringify(row.evidence, null, 2)}
    -            
    - )} -
    - )} -
  • - ); -} - -function ChecklistSummary({ report }: { report: ProviderCertification }) { - return ( -
      - {report.checklist.map((goal) => { - const style = STATUS_STYLES[goal.status]; - const Icon = style.icon; - return ( -
    • - - {goal.label} -
    • - ); - })} -
    - ); -} - -function CertificationReport({ report }: { report: ProviderCertification }) { - const failing = report.rows.filter((row) => row.status === 'fail').length; - const warning = report.rows.filter((row) => row.status === 'warn').length; - - return ( -
    - -
    - - {report.rows.length} checks · {failing} failed · {warning} warnings - - Generated {new Date(report.generated_at).toLocaleString()} -
    -
      - {report.rows.map((row) => ( - - ))} -
    -
    - ); -} - -function getErrorMessage(error: unknown): string { - if (error instanceof Error) return error.message; - return 'Certification request failed'; -} - -function resultStatus( - result: ModelCertificationResult -): CertificationStatus | 'error' { - if (result.error || !result.report) return 'error'; - if (result.report.rows.some((row) => row.status === 'fail')) return 'fail'; - if (result.report.rows.some((row) => row.status === 'warn')) return 'warn'; - return 'ok'; -} - export function ProviderCertificationDialog({ provider, open, onOpenChange, }: ProviderCertificationDialogProps) { - const [checkCache, setCheckCache] = useState(true); - const [workspaceTab, setWorkspaceTab] = useState<'setup' | 'results'>('setup'); - const [selectedModelIds, setSelectedModelIds] = useState([]); - const [pathModes, setPathModes] = useState>({}); - const [selectedModelPaths, setSelectedModelPaths] = useState< - Record - >({}); - const [results, setResults] = useState([]); - const [currentModel, setCurrentModel] = useState<{ - id: string; - index: number; - total: number; - pathCount: number; - } | null>(null); + const [workspaceTab, setWorkspaceTab] = useState<'setup' | 'results'>( + 'setup' + ); + const [setup, setSetup] = useState( + emptyCertificationSetup + ); + const { results, progress, isPending, run, reset } = + useProviderCertificationRunner(provider.id); const models = useQuery({ queryKey: ['provider-models', provider.id], @@ -241,183 +56,27 @@ export function ProviderCertificationDialog({ enabled: open, }); - const certify = useMutation({ - mutationFn: async ({ - modelRuns, - includeCache, - }: { - modelRuns: ModelRun[]; - includeCache: boolean; - }) => { - const completed: ModelCertificationResult[] = []; - setResults([]); - setWorkspaceTab('results'); - - for (const [index, run] of modelRuns.entries()) { - setCurrentModel({ - id: run.modelId, - index: index + 1, - total: modelRuns.length, - pathCount: run.targets.length, - }); - const batch = await Promise.all( - run.targets.map(async (target): Promise => { - const resultKey = `${run.modelId}::${target.path ?? 'default'}`; - try { - const report = await AdminService.certifyProvider(provider.id, { - model_id: run.modelId, - model_path: target.path, - check_cache: includeCache, - }); - return { - resultKey, - modelId: run.modelId, - pathLabel: target.label, - report, - }; - } catch (error) { - return { - resultKey, - modelId: run.modelId, - pathLabel: target.label, - error: getErrorMessage(error), - }; - } - }) - ); - completed.push(...batch); - setResults([...completed]); - } - - return completed; - }, - onSettled: () => setCurrentModel(null), - }); - const { reset: resetCertification } = certify; - useEffect(() => { if (!open) { - resetCertification(); + reset(); setWorkspaceTab('setup'); - setSelectedModelIds([]); - setPathModes({}); - setSelectedModelPaths({}); - setResults([]); - setCurrentModel(null); + setSetup(emptyCertificationSetup()); } - }, [open, resetCertification]); + }, [open, reset]); - const configuredOptions: ModelOption[] = - models.data?.db_models.map((model) => ({ - model, - source: 'configured', - })) ?? []; - const discoveredOptions: ModelOption[] = - models.data?.remote_models.map((model) => ({ - model, - source: 'discovered', - })) ?? []; - const allOptions = [...configuredOptions, ...discoveredOptions]; - const namesById = new Map( - allOptions.map(({ model }) => [model.id, model.name || model.id]) - ); + const modelRuns = buildModelRuns(setup, models.data); + const targetCount = countCertificationTargets(modelRuns); + const modelsNeedingPath = getModelsNeedingPath(setup); + const namesById = getCertificationModelNames(models.data); - const pathsForModel = (modelId: string): CertificationPath[] => - (models.data?.certification_paths[modelId] ?? []).filter( - (path) => path.endpoint_tag - ); - - const pathLabel = (path: CertificationPath): string => - path.endpoint_name && path.endpoint_name !== path.endpoint_tag - ? `${path.endpoint_name} (${path.endpoint_tag})` - : path.endpoint_tag || 'Provider default'; - - const toggleModel = (modelId: string) => { - const isSelected = selectedModelIds.includes(modelId); - setSelectedModelIds((current) => - isSelected - ? current.filter((id) => id !== modelId) - : [...current, modelId] - ); - setPathModes((current) => { - const next = { ...current }; - if (isSelected) delete next[modelId]; - else next[modelId] = 'default'; - return next; - }); - setSelectedModelPaths((current) => { - const next = { ...current }; - if (isSelected) delete next[modelId]; - return next; - }); - resetCertification(); - setResults([]); + const updateSetup = (nextSetup: ProviderCertificationSetup) => { + setSetup(nextSetup); + reset(); }; - const modelsNeedingPath = selectedModelIds.filter( - (modelId) => - pathModes[modelId] === 'selected' && - (selectedModelPaths[modelId]?.length ?? 0) === 0 - ); - - const buildModelRuns = (): ModelRun[] => - selectedModelIds.map((modelId) => { - const paths = pathsForModel(modelId); - const mode = pathModes[modelId] ?? 'default'; - if (mode === 'all') { - return { - modelId, - targets: paths.map((path) => ({ - path: path.path, - label: pathLabel(path), - })), - }; - } - if (mode === 'selected') { - const selected = new Set(selectedModelPaths[modelId] ?? []); - return { - modelId, - targets: paths - .filter((path) => selected.has(path.path)) - .map((path) => ({ path: path.path, label: pathLabel(path) })), - }; - } - return { modelId, targets: [{ label: 'Provider default' }] }; - }); - - const modelRuns = buildModelRuns(); - const targetCount = modelRuns.reduce( - (total, run) => total + run.targets.length, - 0 - ); - - const renderModelGroup = (label: string, options: ModelOption[]) => { - if (options.length === 0) return null; - return ( - - {options.map(({ model }) => ( - toggleModel(model.id)} - disabled={certify.isPending} - > - - ))} - - ); + const runCertification = () => { + setWorkspaceTab('results'); + void run(modelRuns, setup.checkCache); }; return ( @@ -426,9 +85,9 @@ export function ProviderCertificationDialog({ Certify upstream models - Select models, then use the provider default, choose specific - paths, or test every path. Models run one at a time; paths for the - same model run in parallel against {provider.base_url}. + Select models, then use the provider default, choose specific paths, + or test every path. Models run one at a time; paths for the same + model run in parallel against {provider.base_url}. @@ -443,10 +102,10 @@ export function ProviderCertificationDialog({ Setup Results - {(certify.isPending || results.length > 0) && ( + {(isPending || results.length > 0) && ( {results.length}/{targetCount} @@ -458,214 +117,37 @@ export function ProviderCertificationDialog({ value='setup' className='mt-0 min-h-0 overflow-hidden data-[state=active]:flex data-[state=active]:flex-col' > -
    -
    -
    - -
    - - {selectedModelIds.length} selected - - {selectedModelIds.length > 0 && !certify.isPending && ( - - )} -
    -
    - - - - - {models.isLoading ? 'Loading models…' : 'No models found'} - - {renderModelGroup('Configured models', configuredOptions)} - {renderModelGroup('Discovered models', discoveredOptions)} - - - {models.isError && ( -

    - {getErrorMessage(models.error)} -

    - )} -
    - - {selectedModelIds.map((modelId) => { - const paths = pathsForModel(modelId); - const mode = pathModes[modelId] ?? 'default'; - const selectedPaths = selectedModelPaths[modelId] ?? []; - return ( -
    -
    -
    - {namesById.get(modelId) ?? modelId} -
    -
    - {modelId} -
    -
    - {paths.length > 0 ? ( - <> - { - if (!value) return; - setPathModes((current) => ({ - ...current, - [modelId]: value as ModelPathMode, - })); - }} - disabled={certify.isPending} - className='w-full justify-start' - > - Default - - Choose paths - - All paths - - {mode === 'default' && ( -

    - Uses the upstream provider's normal model routing. -

    - )} - {mode === 'selected' && ( - - - - - -
    - {paths.map((path) => { - const checked = selectedPaths.includes( - path.path - ); - return ( - - ); - })} -
    -
    -
    - )} - {mode === 'all' && ( -

    - All {paths.length} paths will run in parallel. -

    - )} - - ) : ( -

    - Only the provider default route is available. -

    - )} -
    - ); - })} - {modelsNeedingPath.length > 0 && ( -

    - Choose at least one path for each model using “Choose paths”. -

    - )} +
    +
    -
    -
    - setCheckCache(value === true)} - disabled={certify.isPending} - /> - -
    +
    + {showEvidence && ( +
    +              {JSON.stringify(row.evidence, null, 2)}
    +            
    + )} +
    + )} + + ); +} + +function CertificationReport({ report }: { report: ProviderCertification }) { + const failing = report.rows.filter((row) => row.status === 'fail').length; + const warning = report.rows.filter((row) => row.status === 'warn').length; + + return ( +
    +
      + {report.checklist.map((goal) => { + const style = STATUS_STYLES[goal.status]; + const Icon = style.icon; + return ( +
    • + + {goal.label} +
    • + ); + })} +
    +
    + + {report.rows.length} checks · {failing} failed · {warning} warnings + + Generated {new Date(report.generated_at).toLocaleString()} +
    +
      + {report.rows.map((row) => ( + + ))} +
    +
    + ); +} + +interface ProviderCertificationResultsProps { + results: ModelCertificationResult[]; + progress?: CertificationProgress | null; + namesById: Map; + emptyMessage?: string; +} + +export function ProviderCertificationResults({ + results, + progress, + namesById, + emptyMessage = 'Results will appear here as certification completes.', +}: ProviderCertificationResultsProps) { + const [activeResult, setActiveResult] = useState(''); + + useEffect(() => { + if ( + results.length > 0 && + !results.some((result) => result.resultKey === activeResult) + ) { + setActiveResult(results[0].resultKey); + } + }, [activeResult, results]); + + return ( +
    + {progress && ( +
    + + + Probing {namesById.get(progress.modelId) ?? progress.modelId} + {progress.pathCount > 1 + ? ` across ${progress.pathCount} paths in parallel` + : ''}{' '} + — model {progress.modelIndex} of {progress.modelTotal} + +
    + )} + + {results.length === 0 ? ( +
    + {emptyMessage} +
    + ) : ( + +
    + + {results.map((result, index) => { + const status = getCertificationResultStatus(result); + const Icon = + status === 'error' ? XCircle : STATUS_STYLES[status].icon; + const modelRouteNumber = results + .slice(0, index + 1) + .filter((item) => item.modelId === result.modelId).length; + const modelRouteCount = results.filter( + (item) => item.modelId === result.modelId + ).length; + return ( + + + + {namesById.get(result.modelId) ?? result.modelId} + + {modelRouteCount > 1 && ( + + {modelRouteNumber} + + )} + + ); + })} + +
    + +
    + {results.map((result) => { + const status = getCertificationResultStatus(result); + return ( + +
    +
    +
    + {namesById.get(result.modelId) ?? result.modelId} +
    + +
    +
    + Model path +
    + {result.pathLabel} +
    +
    +
    + {result.report ? ( + + ) : ( +
    + {result.error ?? 'Certification failed'} +
    + )} +
    + ); + })} +
    +
    + )} +
    + ); +} diff --git a/ui/components/provider-certification-setup.tsx b/ui/components/provider-certification-setup.tsx new file mode 100644 index 00000000..2baa5415 --- /dev/null +++ b/ui/components/provider-certification-setup.tsx @@ -0,0 +1,309 @@ +'use client'; + +import { ChevronDown } from 'lucide-react'; + +import { Button } from '@/components/ui/button'; +import { Checkbox } from '@/components/ui/checkbox'; +import { + Command, + CommandEmpty, + CommandGroup, + CommandInput, + CommandItem, + CommandList, +} from '@/components/ui/command'; +import { Label } from '@/components/ui/label'; +import { + Popover, + PopoverContent, + PopoverTrigger, +} from '@/components/ui/popover'; +import { ToggleGroup, ToggleGroupItem } from '@/components/ui/toggle-group'; +import type { AdminModel, ProviderModels } from '@/lib/api/services/admin'; +import type { + ModelPathMode, + ProviderCertificationSetup, +} from '@/lib/provider-certification'; +import { + emptyCertificationSetup, + getCertificationPathLabel, + getErrorMessage, + getExactCertificationPaths, + getModelsNeedingPath, +} from '@/lib/provider-certification'; + +interface ModelOption { + model: AdminModel; + source: 'configured' | 'discovered'; +} + +interface ProviderCertificationSetupProps { + models?: ProviderModels; + isLoading?: boolean; + error?: unknown; + setup: ProviderCertificationSetup; + onChange: (setup: ProviderCertificationSetup) => void; + disabled?: boolean; + idPrefix: string; +} + +export function getCertificationModelNames( + models: ProviderModels | undefined +): Map { + return new Map( + [...(models?.db_models ?? []), ...(models?.remote_models ?? [])].map( + (model) => [model.id, model.name || model.id] + ) + ); +} + +export function ProviderCertificationSetupPanel({ + models, + isLoading = false, + error, + setup, + onChange, + disabled = false, + idPrefix, +}: ProviderCertificationSetupProps) { + const configuredOptions: ModelOption[] = (models?.db_models ?? []).map( + (model) => ({ model, source: 'configured' }) + ); + const discoveredOptions: ModelOption[] = (models?.remote_models ?? []).map( + (model) => ({ model, source: 'discovered' }) + ); + const namesById = getCertificationModelNames(models); + const modelsNeedingPath = getModelsNeedingPath(setup); + + const toggleModel = (modelId: string) => { + const isSelected = setup.selectedModelIds.includes(modelId); + const pathModes = { ...setup.pathModes }; + const selectedModelPaths = { ...setup.selectedModelPaths }; + if (isSelected) { + delete pathModes[modelId]; + delete selectedModelPaths[modelId]; + } else { + pathModes[modelId] = 'default'; + } + onChange({ + ...setup, + selectedModelIds: isSelected + ? setup.selectedModelIds.filter((id) => id !== modelId) + : [...setup.selectedModelIds, modelId], + pathModes, + selectedModelPaths, + }); + }; + + const renderModelGroup = (label: string, options: ModelOption[]) => { + if (options.length === 0) return null; + return ( + + {options.map(({ model }) => ( + toggleModel(model.id)} + disabled={disabled} + > + + ))} + + ); + }; + + return ( +
    +
    +
    + +
    + + {setup.selectedModelIds.length} selected + + {setup.selectedModelIds.length > 0 && !disabled && ( + + )} +
    +
    + + + + + {isLoading ? 'Loading models…' : 'No models found'} + + {renderModelGroup('Configured models', configuredOptions)} + {renderModelGroup('Discovered models', discoveredOptions)} + + + {Boolean(error) && ( +

    {getErrorMessage(error)}

    + )} +
    + + {setup.selectedModelIds.map((modelId) => { + const paths = getExactCertificationPaths(models, modelId); + const mode = setup.pathModes[modelId] ?? 'default'; + const selectedPaths = setup.selectedModelPaths[modelId] ?? []; + return ( +
    +
    +
    + {namesById.get(modelId) ?? modelId} +
    +
    + {modelId} +
    +
    + {paths.length > 0 ? ( + <> + { + if (!value) return; + onChange({ + ...setup, + pathModes: { + ...setup.pathModes, + [modelId]: value as ModelPathMode, + }, + }); + }} + disabled={disabled} + className='w-full justify-start' + > + Default + + Choose paths + + All paths + + {mode === 'default' && ( +

    + Uses the upstream provider's normal model routing. +

    + )} + {mode === 'selected' && ( + + + + + +
    + {paths.map((path) => { + const checked = selectedPaths.includes(path.path); + return ( + + ); + })} +
    +
    +
    + )} + {mode === 'all' && ( +

    + All {paths.length} paths will run in parallel. +

    + )} + + ) : ( +

    + Only the provider default route is available. +

    + )} +
    + ); + })} + + {modelsNeedingPath.length > 0 && ( +

    + Choose at least one path for each model using “Choose paths”. +

    + )} + +
    + + onChange({ ...setup, checkCache: value === true }) + } + disabled={disabled} + /> + +
    +
    + ); +} diff --git a/ui/hooks/use-provider-certification-runner.ts b/ui/hooks/use-provider-certification-runner.ts new file mode 100644 index 00000000..aff3b914 --- /dev/null +++ b/ui/hooks/use-provider-certification-runner.ts @@ -0,0 +1,128 @@ +'use client'; + +import { useCallback, useEffect, useRef, useState } from 'react'; + +import { AdminService } from '@/lib/api/services/admin'; +import type { + CertificationProgress, + ModelCertificationResult, + ModelRun, +} from '@/lib/provider-certification'; +import { getErrorMessage } from '@/lib/provider-certification'; + +interface RunProviderCertificationOptions { + providerId: number; + modelRuns: ModelRun[]; + includeCache: boolean; + onProgress?: (progress: CertificationProgress | null) => void; + onResults?: (results: ModelCertificationResult[]) => void; +} + +export async function runProviderCertification({ + providerId, + modelRuns, + includeCache, + onProgress, + onResults, +}: RunProviderCertificationOptions): Promise { + const completed: ModelCertificationResult[] = []; + + for (const [index, run] of modelRuns.entries()) { + onProgress?.({ + modelId: run.modelId, + modelIndex: index + 1, + modelTotal: modelRuns.length, + pathCount: run.targets.length, + }); + const batch = await Promise.all( + run.targets.map(async (target): Promise => { + const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`; + try { + const report = await AdminService.certifyProvider(providerId, { + model_id: run.modelId, + model_path: target.path, + check_cache: includeCache, + }); + return { + resultKey, + providerId, + modelId: run.modelId, + pathLabel: target.label, + report, + }; + } catch (error) { + return { + resultKey, + providerId, + modelId: run.modelId, + pathLabel: target.label, + error: getErrorMessage(error), + }; + } + }) + ); + completed.push(...batch); + onResults?.([...completed]); + } + + onProgress?.(null); + return completed; +} + +export function useProviderCertificationRunner(providerId: number) { + const [results, setResults] = useState([]); + const [progress, setProgress] = useState(null); + const [isPending, setIsPending] = useState(false); + const generation = useRef(0); + const mounted = useRef(true); + + useEffect(() => { + mounted.current = true; + return () => { + mounted.current = false; + generation.current += 1; + }; + }, []); + + const reset = useCallback(() => { + generation.current += 1; + setResults([]); + setProgress(null); + setIsPending(false); + }, []); + + const run = useCallback( + async (modelRuns: ModelRun[], includeCache: boolean) => { + const runGeneration = generation.current + 1; + generation.current = runGeneration; + setResults([]); + setProgress(null); + setIsPending(true); + try { + return await runProviderCertification({ + providerId, + modelRuns, + includeCache, + onProgress: (nextProgress) => { + if (mounted.current && generation.current === runGeneration) { + setProgress(nextProgress); + } + }, + onResults: (nextResults) => { + if (mounted.current && generation.current === runGeneration) { + setResults(nextResults); + } + }, + }); + } finally { + if (mounted.current && generation.current === runGeneration) { + setProgress(null); + setIsPending(false); + } + } + }, + [providerId] + ); + + return { results, progress, isPending, run, reset }; +} diff --git a/ui/lib/provider-certification.ts b/ui/lib/provider-certification.ts new file mode 100644 index 00000000..8cdcebc1 --- /dev/null +++ b/ui/lib/provider-certification.ts @@ -0,0 +1,116 @@ +import type { + CertificationPath, + CertificationStatus, + ProviderCertification, + ProviderModels, +} from '@/lib/api/services/admin'; + +export type ModelPathMode = 'default' | 'selected' | 'all'; + +export interface ProviderCertificationSetup { + selectedModelIds: string[]; + pathModes: Record; + selectedModelPaths: Record; + checkCache: boolean; +} + +export interface ModelPathTarget { + path?: string; + label: string; +} + +export interface ModelRun { + modelId: string; + targets: ModelPathTarget[]; +} + +export interface CertificationProgress { + modelId: string; + modelIndex: number; + modelTotal: number; + pathCount: number; +} + +export interface ModelCertificationResult { + resultKey: string; + providerId: number; + modelId: string; + pathLabel: string; + report?: ProviderCertification; + error?: string; +} + +export const emptyCertificationSetup = (): ProviderCertificationSetup => ({ + selectedModelIds: [], + pathModes: {}, + selectedModelPaths: {}, + checkCache: true, +}); + +export const getExactCertificationPaths = ( + models: ProviderModels | undefined, + modelId: string +): CertificationPath[] => + (models?.certification_paths[modelId] ?? []).filter( + (path) => path.endpoint_tag + ); + +export const getCertificationPathLabel = (path: CertificationPath): string => + path.endpoint_name && path.endpoint_name !== path.endpoint_tag + ? `${path.endpoint_name} (${path.endpoint_tag})` + : path.endpoint_tag || 'Provider default'; + +export const getModelsNeedingPath = ( + setup: ProviderCertificationSetup +): string[] => + setup.selectedModelIds.filter( + (modelId) => + setup.pathModes[modelId] === 'selected' && + (setup.selectedModelPaths[modelId]?.length ?? 0) === 0 + ); + +export const buildModelRuns = ( + setup: ProviderCertificationSetup, + models: ProviderModels | undefined +): ModelRun[] => + setup.selectedModelIds.map((modelId) => { + const paths = getExactCertificationPaths(models, modelId); + const mode = setup.pathModes[modelId] ?? 'default'; + if (mode === 'all') { + return { + modelId, + targets: paths.map((path) => ({ + path: path.path, + label: getCertificationPathLabel(path), + })), + }; + } + if (mode === 'selected') { + const selected = new Set(setup.selectedModelPaths[modelId] ?? []); + return { + modelId, + targets: paths + .filter((path) => selected.has(path.path)) + .map((path) => ({ + path: path.path, + label: getCertificationPathLabel(path), + })), + }; + } + return { modelId, targets: [{ label: 'Provider default' }] }; + }); + +export const countCertificationTargets = (runs: ModelRun[]): number => + runs.reduce((total, run) => total + run.targets.length, 0); + +export const getCertificationResultStatus = ( + result: ModelCertificationResult +): CertificationStatus | 'error' => { + if (result.error || !result.report) return 'error'; + if (result.report.rows.some((row) => row.status === 'fail')) return 'fail'; + if (result.report.rows.some((row) => row.status === 'warn')) return 'warn'; + return 'ok'; +}; + +export const getErrorMessage = (error: unknown): string => + error instanceof Error ? error.message : 'Certification request failed'; From 0b772dba49c8fcd3abe6c09d5ef5a98f24c46404 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 00:57:24 +0200 Subject: [PATCH 11/18] fix: satisfy mypy list variance in certification cache tests and prettier in admin schema --- tests/unit/test_certification_cache.py | 4 ++-- ui/lib/api/services/admin.ts | 5 +---- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py index 2ca56482..111659ba 100644 --- a/tests/unit/test_certification_cache.py +++ b/tests/unit/test_certification_cache.py @@ -257,7 +257,7 @@ class TestCostMarginRow: cache_read=4e-9, sats_usd=sats_usd, ) - payloads = [ + payloads: list[dict[str, Any] | None] = [ _payload( { "prompt_tokens": 31, @@ -306,7 +306,7 @@ class TestCostMarginRow: cache_read=4e-9, sats_usd=sats_usd, ) - payloads = [ + payloads: list[dict[str, Any] | None] = [ _payload( { "prompt_tokens": 31, diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index ba665da0..0ae0e277 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -135,10 +135,7 @@ export const ProviderModelsSchema = z.object({ }), db_models: z.array(AdminModelSchema), remote_models: z.array(AdminModelSchema), - certification_paths: z.record( - z.string(), - z.array(CertificationPathSchema) - ), + certification_paths: z.record(z.string(), z.array(CertificationPathSchema)), }); export type ProviderType = z.infer; From 41b624c43669a5cae1465c218c050fea2e84c21c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 22:26:03 +0200 Subject: [PATCH 12/18] fix: publish resolved sats price, honour fixed pricing and guard certify endpoint --- routstr/core/admin.py | 16 ++- routstr/upstream/certification.py | 87 +++++++++---- routstr/upstream/certification_cache.py | 18 ++- tests/integration/test_certify_endpoint.py | 20 ++- tests/unit/test_certification.py | 121 ++++++++++++++++++ tests/unit/test_certification_cache.py | 26 ++++ ui/hooks/use-provider-certification-runner.ts | 55 ++++---- ui/lib/provider-certification.ts | 3 +- 8 files changed, 283 insertions(+), 63 deletions(-) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index bd5b4e18..745bb8e4 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1549,7 +1549,6 @@ async def certify_upstream_provider( :mod:`routstr.upstream.certification`, and a ``checklist`` of the operator-facing goals. """ - from ..payment.price import sats_usd_price from ..upstream.certification import ( MAX_PROBE_TIMEOUT_SECONDS, PROBE_TIMEOUT_SECONDS, @@ -1707,7 +1706,16 @@ async def certify_upstream_provider( *skipped_cache_rows("Skipped — no model to probe."), ] else: - sats_to_usd = sats_usd_price() + from ..payment import price as price_module + + # Never fetch the price inline: the lifespan task owns it, and a fetch + # here could block the request for the exchange timeout. + sats_to_usd = price_module.SATS_USD_PRICE + if not sats_to_usd: + raise HTTPException( + status_code=503, + detail="sats/USD price is not initialized yet; retry shortly", + ) if selected_path is not None: from ..upstream.model_paths import apply_model_path_pricing @@ -1717,8 +1725,8 @@ async def certify_upstream_provider( provider.provider_fee, sats_to_usd, ) - # Clamp the admin-supplied timeout so a probe cannot hold the request - # open indefinitely. + # Clamp the admin-supplied timeout per upstream call. The run makes up + # to five calls, so the request can stay open for up to five times it. requested = ( payload.timeout_seconds if payload.timeout_seconds is not None diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 5a87c8d0..188bd16b 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -515,6 +515,42 @@ def _reported_usd_cost(payload: dict[str, Any]) -> float: return 0.0 +def _fixed_token_pricing_active() -> bool: + """Whether node-wide fixed per-1k pricing overrides the model's rates.""" + from ..core.settings import settings + + return bool( + settings.fixed_pricing + and (settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens) + ) + + +def _token_rates(sats_pricing: Any) -> tuple[float, float, float, float]: + """The msats-per-1k rates the engine bills tokens at. + + Mirrors ``_get_pricing_rates``'s selection: node-wide fixed pricing + overrides the model's own rates, with cache tokens at the input rate. + + Returns ``(input, output, cache_read, cache_write)``. + """ + from ..core.settings import settings + + if _fixed_token_pricing_active(): + fixed_input = float(settings.fixed_per_1k_input_tokens) * 1000.0 + fixed_output = float(settings.fixed_per_1k_output_tokens) * 1000.0 + return fixed_input, fixed_output, fixed_input, fixed_input + + input_rate = float(sats_pricing.prompt) * 1_000_000.0 + output_rate = float(sats_pricing.completion) * 1_000_000.0 + cache_read_rate = ( + float(sats_pricing.input_cache_read or 0.0) * 1_000_000.0 or input_rate + ) + cache_write_rate = ( + float(sats_pricing.input_cache_write or 0.0) * 1_000_000.0 or input_rate + ) + return input_rate, output_rate, cache_read_rate, cache_write_rate + + def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int]: """Re-derive the token-priced charge independently of the engine. @@ -525,18 +561,10 @@ def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int] Returns ``(total_msats, input_msats, output_msats)``. Raises ``ValueError`` on a non-finite rate, which would otherwise crash ``math.ceil`` downstream. """ - input_rate = float(sats_pricing.prompt) * 1_000_000.0 - output_rate = float(sats_pricing.completion) * 1_000_000.0 - cache_read_rate = ( - float(sats_pricing.input_cache_read or 0.0) * 1_000_000.0 or input_rate - ) - cache_write_rate = ( - 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) + rates = _token_rates(sats_pricing) if not all(math.isfinite(rate) for rate in rates): raise ValueError(f"non-finite pricing rate in {rates!r}") + input_rate, output_rate, cache_read_rate, cache_write_rate = rates calc_input = round(usage.input_tokens / 1000 * input_rate, 3) calc_output = round(usage.output_tokens / 1000 * output_rate, 3) @@ -587,6 +615,18 @@ def cost_prompt_completion_row( "sats_usd_price": sats_to_usd, } + # Checked before the engine's error: with no price the engine cannot + # succeed, and that is a gap in the run's inputs, not a node fault. + if not pricing_known: + return certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + "No pricing is known for this model, so the charge cannot be " + "verified. Configure the model on the node, or pass explicit " + "prices, to certify this row.", + evidence, + ) if isinstance(cost_data, CostDataError): evidence["error"] = cost_data.message return certification_row( @@ -613,16 +653,6 @@ def cost_prompt_completion_row( "verify the charge against.", evidence, ) - if not pricing_known: - return certification_row( - "cost.prompt_completion", - STATUS_WARN, - "Prompt and completion cost calculated", - "No pricing is known for this model, so the charge cannot be " - "verified. Configure the model on the node, or pass explicit " - "prices, to certify this row.", - evidence, - ) reported_usd = _reported_usd_cost(payload) try: @@ -865,15 +895,24 @@ async def _resolve_sats_usd_price(override: float | None) -> float | None: 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 + # The cost engine reads the module globals rather than this return value, + # so a resolved price is published there too or every token-priced cost + # row fails on "SATS price not initialized". + if override is not None: + if not (math.isfinite(override) and override > 0): + return None + price_module.SATS_USD_PRICE = override + price_module.BTC_USD_PRICE = override * price_module.SATS_PER_BTC + return override + 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 + sats_price = float(price_module.BTC_USD_PRICE) / price_module.SATS_PER_BTC + price_module.SATS_USD_PRICE = sats_price + return sats_price try: await price_module._update_prices() diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index 3ee74e7a..d8f21f13 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -32,7 +32,9 @@ from .certification import ( STATUS_WARN, _expected_token_msats, _expected_usd_msats, + _fixed_token_pricing_active, _reported_usd_cost, + _token_rates, certification_row, safe_row, ) @@ -341,8 +343,7 @@ def cache_billing_row( ) pricing = model.sats_pricing - cache_read_rate = float(pricing.input_cache_read or 0.0) - input_rate = float(pricing.prompt) + input_rate, _, cache_read_rate, _ = _token_rates(pricing) full_usage = NormalizedUsage( input_tokens=usage.input_tokens + usage.cache_read_tokens @@ -367,8 +368,8 @@ def cache_billing_row( evidence.update( { "usage": usage.dict(), - "cache_read_rate_sats": cache_read_rate, - "input_rate_sats": input_rate, + "cache_read_rate_msats_per_1k": cache_read_rate, + "input_rate_msats_per_1k": input_rate, "actual_total_msats": actual_total, "expected_total_msats": expected_total, "full_price_total_msats": full_total, @@ -396,13 +397,18 @@ def cache_billing_row( evidence, ) if cache_read_rate <= 0.0 or cache_read_rate >= input_rate: + reason = ( + "the node uses fixed per-1k pricing" + if _fixed_token_pricing_active() + else "no discounted cache-read rate is configured" + ) return certification_row( ROW_BILLING, STATUS_WARN, TITLE_BILLING, f"Cached reads are billed at the full input rate ({actual_total} " - "msats) because no discounted cache-read rate is configured; " - "clients pay more than the upstream charges.", + f"msats) because {reason}; clients pay more than the upstream " + "charges.", evidence, ) return certification_row( diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index c7aa9151..ac38a9e0 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -399,7 +399,7 @@ async def test_certify_uses_selected_path_pricing_for_margin( respx.post(f"{base_url}/chat/completions").mock(side_effect=_respond) sats_usd = 0.0008616302499999999 - with patch("routstr.payment.price.sats_usd_price", return_value=sats_usd): + with patch("routstr.payment.price.SATS_USD_PRICE", sats_usd): resp = await integration_client.post( f"/admin/api/upstream-providers/{provider_id}/certify", headers=_admin_headers(), @@ -416,6 +416,24 @@ async def test_certify_uses_selected_path_pricing_for_margin( assert "15 < 26" in margin["detail"] +@pytest.mark.integration +@pytest.mark.asyncio +async def test_certify_returns_503_when_price_uninitialized( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + with patch("routstr.payment.price.SATS_USD_PRICE", None): + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + + assert resp.status_code == 503, resp.text + assert "sats/USD price is not initialized" in resp.json()["detail"] + + @pytest.mark.integration @pytest.mark.asyncio @respx.mock diff --git a/tests/unit/test_certification.py b/tests/unit/test_certification.py index 1e963a79..8c1d82e4 100644 --- a/tests/unit/test_certification.py +++ b/tests/unit/test_certification.py @@ -10,6 +10,8 @@ from __future__ import annotations from typing import Any +import pytest + from routstr.upstream.certification import ( STATUS_FAIL, STATUS_OK, @@ -606,3 +608,122 @@ class TestBuildChecklist: goals = {item["goal"]: item["status"] for item in checklist} assert goals["heartbeat"] == STATUS_FAIL assert goals["usage_data"] == STATUS_WARN + + +# --- engine-consistent pricing ---------------------------------------------- + + +class TestStandaloneCostRow: + """The standalone runner prices a real completion through the engine.""" + + @staticmethod + def _client() -> Any: + import httpx + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "cert-model"}]}) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-1", + "object": "chat.completion", + "model": "cert-model", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + }, + ) + + return httpx.AsyncClient(transport=httpx.MockTransport(handler)) + + @staticmethod + def _cost_row(result: dict[str, Any]) -> dict[str, Any]: + return next(r for r in result["rows"] if r["id"] == "cost.prompt_completion") + + async def test_override_price_reaches_the_engine( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.payment import price as price_module + from routstr.upstream.certification import certify_upstream_url + + monkeypatch.setattr(price_module, "SATS_USD_PRICE", None) + monkeypatch.setattr(price_module, "BTC_USD_PRICE", None) + + async with self._client() as client: + result = await certify_upstream_url( + "https://upstream.example/v1", + model_id="cert-model", + sats_usd_price=5e-7, + prompt_price=1e-6, + completion_price=2e-6, + client=client, + check_cache=False, + ) + + row = self._cost_row(result) + assert row["status"] == STATUS_OK, row["detail"] + + async def test_no_price_available_warns( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.payment import price as price_module + from routstr.upstream.certification import certify_upstream_url + + async def offline() -> None: + raise RuntimeError("exchange feed unreachable") + + monkeypatch.setattr(price_module, "SATS_USD_PRICE", None) + monkeypatch.setattr(price_module, "BTC_USD_PRICE", None) + monkeypatch.setattr(price_module, "_update_prices", offline) + + async with self._client() as client: + result = await certify_upstream_url( + "https://upstream.example/v1", + model_id="cert-model", + prompt_price=1e-6, + completion_price=2e-6, + client=client, + check_cache=False, + ) + + row = self._cost_row(result) + assert row["status"] == STATUS_WARN, row["detail"] + + +class TestFixedPricingCostRow: + async def test_fixed_pricing_node_certifies_ok( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.core.settings import settings + from routstr.payment import price as price_module + from routstr.payment.cost_calculation import calculate_cost + + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 3) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 7) + monkeypatch.setattr(price_module, "SATS_USD_PRICE", 0.0005) + + model = TestCostPromptCompletion()._model() + payload = { + "model": "test-model", + "usage": {"prompt_tokens": 1000, "completion_tokens": 500}, + } + cost_data = await calculate_cost( + payload, 1_000_000, model_obj=model, provider_fee=1.0 + ) + + row = cost_prompt_completion_row( + model=model, + probe=_probe(chat_status=200, chat_payload=payload), + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_OK, row["detail"] + assert row["evidence"]["expected_total_msats"] == 3000 + 3500 diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py index 111659ba..2d87a8cf 100644 --- a/tests/unit/test_certification_cache.py +++ b/tests/unit/test_certification_cache.py @@ -463,6 +463,32 @@ class TestErrorBranches: assert row["status"] == STATUS_FAIL assert "error" in row["evidence"] + @pytest.mark.asyncio + async def test_billing_warns_under_fixed_pricing( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.core.settings import settings + from routstr.payment.cost_calculation import calculate_cost + + fixed_in, fixed_out = 2.0, 3.0 + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", fixed_in) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", fixed_out) + monkeypatch.setattr("routstr.payment.price.SATS_USD_PRICE", SATS_USD) + + model = _model(cache_read=1.4e-8) + payload = _payload(CACHED) + cost = await calculate_cost(payload, 10**9, model_obj=model, provider_fee=1.0) + row = cache_billing_row( + model=model, + probe=_probe([_payload(UNCACHED), payload]), + cost_data=cost, + ) + assert row["status"] == STATUS_WARN, row + assert "fixed per-1k pricing" in row["detail"] + assert row["evidence"]["input_rate_msats_per_1k"] == fixed_in * 1000 + assert row["evidence"]["cache_read_rate_msats_per_1k"] == fixed_in * 1000 + def test_margin_fails_on_zero_sats_price(self) -> None: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) row = cost_margin_row( diff --git a/ui/hooks/use-provider-certification-runner.ts b/ui/hooks/use-provider-certification-runner.ts index aff3b914..fb7a9f7d 100644 --- a/ui/hooks/use-provider-certification-runner.ts +++ b/ui/hooks/use-provider-certification-runner.ts @@ -34,33 +34,34 @@ export async function runProviderCertification({ modelTotal: modelRuns.length, pathCount: run.targets.length, }); - const batch = await Promise.all( - run.targets.map(async (target): Promise => { - const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`; - try { - const report = await AdminService.certifyProvider(providerId, { - model_id: run.modelId, - model_path: target.path, - check_cache: includeCache, - }); - return { - resultKey, - providerId, - modelId: run.modelId, - pathLabel: target.label, - report, - }; - } catch (error) { - return { - resultKey, - providerId, - modelId: run.modelId, - pathLabel: target.label, - error: getErrorMessage(error), - }; - } - }) - ); + // Sequential on purpose: each run spends real upstream credits, and + // parallel paths multiply that spend and the admin request load. + const batch: ModelCertificationResult[] = []; + for (const target of run.targets) { + const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`; + try { + const report = await AdminService.certifyProvider(providerId, { + model_id: run.modelId, + model_path: target.path, + check_cache: includeCache, + }); + batch.push({ + resultKey, + providerId, + modelId: run.modelId, + pathLabel: target.label, + report, + }); + } catch (error) { + batch.push({ + resultKey, + providerId, + modelId: run.modelId, + pathLabel: target.label, + error: getErrorMessage(error), + }); + } + } completed.push(...batch); onResults?.([...completed]); } diff --git a/ui/lib/provider-certification.ts b/ui/lib/provider-certification.ts index 8cdcebc1..1b9a0d20 100644 --- a/ui/lib/provider-certification.ts +++ b/ui/lib/provider-certification.ts @@ -1,3 +1,4 @@ +import { getApiErrorMessage } from '@/lib/api/errors'; import type { CertificationPath, CertificationStatus, @@ -113,4 +114,4 @@ export const getCertificationResultStatus = ( }; export const getErrorMessage = (error: unknown): string => - error instanceof Error ? error.message : 'Certification request failed'; + getApiErrorMessage(error, 'Certification request failed'); From 943aa8f6fa7e5c1f038682538a3083ad4dbfe9a8 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 01:37:48 +0200 Subject: [PATCH 13/18] fix: resolve upstream certification review findings --- routstr/core/admin.py | 33 +- routstr/upstream/certification.py | 42 ++- routstr/upstream/certification_cache.py | 25 +- tests/integration/test_certify_alias_paths.py | 130 +++++++ .../test_certification_review_regressions.py | 308 ++++++++++++++++ ui/app/providers/certification/page.tsx | 64 +++- .../provider-certification-dialog.tsx | 12 +- ...use-provider-certification-runner.test.mjs | 328 ++++++++++++++++++ ui/hooks/use-provider-certification-runner.ts | 18 +- ui/lib/provider-certification.ts | 6 + 10 files changed, 913 insertions(+), 53 deletions(-) create mode 100644 tests/integration/test_certify_alias_paths.py create mode 100644 tests/unit/test_certification_review_regressions.py create mode 100644 ui/hooks/use-provider-certification-runner.test.mjs diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 745bb8e4..bf3aac15 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1236,9 +1236,7 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: ] path_result = await session.exec( - select(ModelPathRow).where( - ModelPathRow.upstream_provider_id == provider_pk - ) + select(ModelPathRow).where(ModelPathRow.upstream_provider_id == provider_pk) ) path_rows = list(path_result.all()) paths_by_public_id: dict[str, list[dict[str, object]]] = {} @@ -1251,12 +1249,11 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: } ) - from ..upstream.model_paths import public_model_id + from ..upstream.model_paths import exposed_model_id certification_paths: dict[str, list[dict[str, object]]] = {} for model in [*db_models, *filtered_remote_models]: - forwarded_id = model.forwarded_model_id or model.id - paths = paths_by_public_id.get(public_model_id(forwarded_id).lower(), []) + paths = paths_by_public_id.get(exposed_model_id(model).lower(), []) certification_paths[model.id] = paths return { @@ -1570,19 +1567,11 @@ async def certify_upstream_provider( endpoint_tag: str | None = None selected_path: ModelPathRow | None = None if payload.model_path is not None: - from ..proxy import _model_ids_match from ..upstream.model_paths import decode_model_path selector = decode_model_path(payload.model_path) if selector is None: raise HTTPException(status_code=400, detail="Malformed model path") - if payload.model_id is None or not _model_ids_match( - payload.model_id, selector.model_id - ): - raise HTTPException( - status_code=400, - detail="Model path does not match the selected model", - ) path_result = await session.exec( select(ModelPathRow).where( ModelPathRow.upstream_provider_id == provider_pk, @@ -1652,13 +1641,25 @@ async def certify_upstream_provider( ( model for model in upstream.get_cached_models() - if model.id == model_id - or model.forwarded_model_id == model_id + if model.id == model_id or model.forwarded_model_id == model_id ), None, ) if model_obj is not None: break + if selected_path is not None: + from ..upstream.model_paths import exposed_model_id + + selected_id = exposed_model_id(model_obj) if model_obj else model_id + if ( + payload.model_id is None + or selected_id is None + or selected_id.lower() != selector.model_id.lower() + ): + raise HTTPException( + status_code=400, + detail="Model path does not match the selected model", + ) if model_obj is None: from ..upstream.certification import ( STATUS_WARN, diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 188bd16b..1051b019 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -192,8 +192,8 @@ async def probe_upstream( ) -> ProbeResult: """Call the upstream's ``/models`` and a one-token completion. - A transport failure on either call is recorded on the result rather - than raised: a dead upstream is a ``fail`` row, not a failed request. + Each HTTP call, including its body read, has an elapsed-time deadline. + A transport failure is a ``fail`` row, not a failed admin request. """ base = base_url.rstrip("/") result = ProbeResult( @@ -213,7 +213,8 @@ async def probe_upstream( try: started = time.monotonic() try: - response = await client.get(result.models_url, headers=headers) + async with asyncio.timeout(timeout): + response = await client.get(result.models_url, headers=headers) result.models_status = response.status_code result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) try: @@ -246,9 +247,10 @@ async def probe_upstream( "allow_fallbacks": False, } try: - response = await client.post( - result.chat_url, json=request_body, headers=headers - ) + async with asyncio.timeout(timeout): + response = await client.post( + result.chat_url, json=request_body, headers=headers + ) result.chat_status = response.status_code result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) try: @@ -507,6 +509,9 @@ def _reported_usd_cost(payload: dict[str, Any]) -> float: total = coerce_rate(cost_details.get("total_cost")) if total is not None and total > 0: return total + inference = coerce_rate(cost_details.get("upstream_inference_cost")) + if inference is not None and inference > 0 and usage.get("is_byok"): + return inference + (coerce_rate(usage.get("cost")) or 0.0) for source in (usage, payload): for field in ("total_cost", "cost"): value = coerce_rate(source.get(field)) @@ -830,7 +835,11 @@ async def run_live_checks( if not check_cache: rows.extend(skipped_cache_rows("Skipped — cache checks disabled.")) - elif probe.chat_payload is None: + elif ( + probe.chat_payload is None + or probe.chat_status is None + or not 200 <= probe.chat_status < 300 + ): rows.extend( skipped_cache_rows("Skipped — the completion probe did not succeed.") ) @@ -928,7 +937,14 @@ async def _resolve_sats_usd_price(override: float | None) -> float | None: def _model_from_usd_pricing( - model_id: str, prompt_usd: float, completion_usd: float, sats_to_usd: float + model_id: str, + prompt_usd: float, + completion_usd: float, + sats_to_usd: float, + *, + provider_fee: float = 1.0, + cache_read_usd: float | None = None, + cache_write_usd: float | None = None, ) -> "Model": """A throwaway ``Model`` carrying just enough to exercise the cost engine.""" from ..payment.models import ( @@ -951,7 +967,12 @@ def _model_from_usd_pricing( tokenizer="unknown", instruct_type=None, ), - pricing=Pricing(prompt=prompt_usd, completion=completion_usd), + pricing=Pricing( + prompt=prompt_usd * provider_fee, + completion=completion_usd * provider_fee, + input_cache_read=(cache_read_usd or 0.0) * provider_fee, + input_cache_write=(cache_write_usd or 0.0) * provider_fee, + ), sats_pricing=None, per_request_limits=None, top_provider=None, @@ -1042,6 +1063,9 @@ async def certify_upstream_url( resolved_prompt or 0.0, resolved_completion or 0.0, sats_to_usd or 1.0, + provider_fee=provider_fee, + cache_read_usd=_as_price(entry.get("cache_read_input_token_cost")), + cache_write_usd=_as_price(entry.get("cache_creation_input_token_cost")), ) rows = await run_live_checks( base_url, diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index d8f21f13..d8ed2dad 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -13,6 +13,7 @@ with ``httpx`` and never enter the billing path. from __future__ import annotations +import asyncio import time from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any @@ -130,10 +131,12 @@ async def _post_completion( url: str, body: dict[str, Any], headers: dict[str, str], + timeout: float, ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: started = time.monotonic() try: - response = await client.post(url, json=body, headers=headers) + async with asyncio.timeout(timeout): + response = await client.post(url, json=body, headers=headers) except Exception as exc: # noqa: BLE001 - transport failure is a row status latency = round((time.monotonic() - started) * 1000, 2) return None, None, f"{type(exc).__name__}: {exc}", latency @@ -180,7 +183,8 @@ async def probe_cache( The first attempt marks the prefix with an Anthropic-style ``cache_control`` part. Upstreams that reject the part get a plain string - retry, and the second call mirrors whichever format succeeded. + retry on HTTP 400/422, and the second call mirrors whichever format + succeeded. Each call's elapsed deadline includes the response body. """ base = base_url.rstrip("/") result = CacheProbeResult( @@ -200,14 +204,16 @@ async def probe_cache( result.chat_url, _request_body(model_id, prefix, "cache_control", endpoint_tag), headers, + timeout, ) - if not _is_2xx(first[0]) and first[0] is not None: + if first[0] in (400, 422): result.request_format = "plain" first = await _post_completion( client, result.chat_url, _request_body(model_id, prefix, "plain", endpoint_tag), headers, + timeout, ) _record(result, first) if not _is_2xx(first[0]): @@ -217,6 +223,7 @@ async def probe_cache( result.chat_url, _request_body(model_id, prefix, result.request_format, endpoint_tag), headers, + timeout, ) _record(result, second) finally: @@ -282,7 +289,17 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]: evidence, ) - raw_keys = _raw_cache_keys(payload.get("usage")) + known_write_fields = { + "cache_creation_input_tokens", + "prompt_tokens_details.cache_creation_tokens", + "prompt_tokens_details.cache_write_tokens", + "input_tokens_details.cache_write_tokens", + } + raw_keys = [ + key + for key in _raw_cache_keys(payload.get("usage")) + if key not in known_write_fields + ] if raw_keys: evidence["unrecognised_cache_fields"] = raw_keys return certification_row( diff --git a/tests/integration/test_certify_alias_paths.py b/tests/integration/test_certify_alias_paths.py new file mode 100644 index 00000000..44e10f7c --- /dev/null +++ b/tests/integration/test_certify_alias_paths.py @@ -0,0 +1,130 @@ +"""Exact certification paths retain the provider model's forwarded identity.""" + +import json +from typing import Any +from unittest.mock import patch + +import pytest +import respx +from httpx import AsyncClient, Response +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ModelPathRow +from routstr.proxy import reinitialize_upstreams +from routstr.upstream.model_paths import encode_model_path +from tests.integration.test_certify_endpoint import ( + _admin_headers, + _make_provider, + _model_row, +) + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("forwarded", ["anthropic/claude-opus-4.6", "remote-id"]) +@pytest.mark.parametrize("enabled", [True, False]) +@respx.mock +async def test_certify_forwarded_alias_listed_path_succeeds( + integration_session: AsyncSession, + integration_client: AsyncClient, + forwarded: str, + enabled: bool, + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr import proxy + + for name in ("_upstreams", "_provider_map", "_unique_models"): + monkeypatch.setattr(proxy, name, getattr(proxy, name).copy()) + base_url = "https://certify-upstream.example/v1" + respx.get(f"{base_url}/models").mock(return_value=Response(200, json={"data": []})) + chat = respx.post(f"{base_url}/chat/completions").mock( + return_value=Response( + 200, + json={ + "model": forwarded, + "usage": {"prompt_tokens": 5, "completion_tokens": 1}, + }, + ) + ) + provider = await _make_provider(integration_session) + model = _model_row(provider.id, model_id="local-alias") # type: ignore[arg-type] + model.forwarded_model_id = forwarded + model.enabled = enabled + model_path = encode_model_path(base_url, forwarded, "endpoint") + integration_session.add(model) + integration_session.add( + ModelPathRow( + upstream_provider_id=provider.id, + model_id=forwarded, + path=model_path, + endpoint_tag="endpoint", + provider_slug="mock", + provider_type="generic", + ) + ) + other_id = f"other/{forwarded.rsplit('/', 1)[-1]}" + other_path = encode_model_path(base_url, other_id, "other-endpoint") + integration_session.add( + ModelPathRow( + upstream_provider_id=provider.id, + model_id=other_id, + path=other_path, + endpoint_tag="other-endpoint", + provider_slug="mock", + provider_type="generic", + ) + ) + await integration_session.commit() + with ( + patch("routstr.payment.models.sats_usd_price", return_value=0.0005), + patch("routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005), + patch("routstr.payment.price.SATS_USD_PRICE", 0.0005), + ): + await reinitialize_upstreams() + listed = await integration_client.get( + f"/admin/api/upstream-providers/{provider.id}/models", + headers=_admin_headers(), + ) + assert listed.status_code == 200, listed.text + assert ( + listed.json()["certification_paths"]["local-alias"][0]["path"] == model_path + ) + response = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={ + "model_id": "local-alias", + "model_path": model_path, + "check_cache": False, + }, + ) + if not enabled: + assert response.status_code == 400, response.text + assert chat.call_count == 0 + return + assert response.status_code == 200, response.text + assert chat.call_count == 1 + body: dict[str, Any] = json.loads(chat.calls[0].request.content) + assert body["model"] == forwarded + assert body["provider"] == {"order": ["endpoint"], "allow_fallbacks": False} + mismatch = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={ + "model_id": "wrong-alias", + "model_path": model_path, + "check_cache": False, + }, + ) + assert mismatch.status_code == 400 + wrong_prefix = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={ + "model_id": "local-alias", + "model_path": other_path, + "check_cache": False, + }, + ) + assert wrong_prefix.status_code == 400 + assert chat.call_count == 1 diff --git a/tests/unit/test_certification_review_regressions.py b/tests/unit/test_certification_review_regressions.py new file mode 100644 index 00000000..81dac641 --- /dev/null +++ b/tests/unit/test_certification_review_regressions.py @@ -0,0 +1,308 @@ +"""Regression coverage for certification pricing and probe lifecycle fixes.""" + +import asyncio +from collections.abc import AsyncIterator +from typing import Any + +import httpx +import pytest + +from routstr.payment import price +from routstr.payment.cost_calculation import calculate_cost +from routstr.upstream.certification import ( + _model_from_usd_pricing, + certify_upstream_url, + cost_prompt_completion_row, + probe_upstream, + run_live_checks, +) +from routstr.upstream.certification_cache import ( + CacheProbeResult, + cache_reported_row, + cost_margin_row, + probe_cache, +) + + +@pytest.mark.asyncio +async def test_byok_cost_and_margin_include_inference_and_routing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(price, "SATS_USD_PRICE", 0.001) + model = _model_from_usd_pricing("test-model", 1e-6, 2e-6, 0.001) + payload = { + "model": "test-model", + "usage": { + "prompt_tokens": 10, + "completion_tokens": 1, + "is_byok": True, + "cost": 0.00005, + "cost_details": {"upstream_inference_cost": 0.001}, + }, + } + from routstr.upstream.certification import ProbeResult + + probe = ProbeResult( + base_url="https://mock.example/v1", + models_url="https://mock.example/v1/models", + chat_url="https://mock.example/v1/chat/completions", + chat_status=200, + chat_payload=payload, + ) + cost = await calculate_cost(payload, 1_000_000, model_obj=model, provider_fee=1) + row = cost_prompt_completion_row( + model=model, probe=probe, cost_data=cost, provider_fee=1, sats_to_usd=0.001 + ) + assert row["status"] == "ok" + assert row["evidence"]["expected_total_msats"] == 1050 + margin = cost_margin_row( + model=model, payloads=[payload], provider_fee=1, sats_to_usd=0.001 + ) + assert margin["status"] == "fail" + assert margin["evidence"]["samples"][0]["upstream_msats_with_fee"] == 1050 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("usage", "entry", "fee", "explicit", "expected"), + [ + ( + {"prompt_tokens": 1000, "completion_tokens": 1}, + {"input_cost_per_token": 9e-6, "output_cost_per_token": 9e-6}, + 2, + True, + 2004, + ), + ( + { + "prompt_tokens": 1000, + "completion_tokens": 1, + "prompt_tokens_details": {"cached_tokens": 900}, + }, + { + "cache_read_input_token_cost": -1, + "cache_creation_input_token_cost": float("inf"), + }, + 1, + False, + 1002, + ), + ( + { + "prompt_tokens": 1000, + "completion_tokens": 1, + "prompt_tokens_details": {"cached_tokens": 900}, + }, + {"cache_read_input_token_cost": 1e-7}, + 1, + False, + 192, + ), + ( + { + "input_tokens": 100, + "output_tokens": 1, + "cache_creation_input_tokens": 900, + }, + {"cache_creation_input_token_cost": 1.25e-6}, + 2, + False, + 2454, + ), + ( + {"prompt_tokens": 1000, "completion_tokens": 1, "cost": 0.001002}, + {}, + 2, + True, + 2004, + ), + ], +) +async def test_standalone_preserves_fee_cache_rates_and_usd_fee( + monkeypatch: pytest.MonkeyPatch, + usage: dict[str, Any], + entry: dict[str, float], + fee: float, + explicit: bool, + expected: int, +) -> None: + from routstr.payment import models + + monkeypatch.setattr(price, "SATS_USD_PRICE", None) + monkeypatch.setattr(price, "BTC_USD_PRICE", None) + monkeypatch.setattr( + models, + "litellm_cost_entry", + lambda _: { + "input_cost_per_token": 1e-6, + "output_cost_per_token": 2e-6, + **entry, + }, + ) + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "test-model"}]}) + return httpx.Response(200, json={"model": "test-model", "usage": usage}) + + kwargs = {"prompt_price": 1e-6, "completion_price": 2e-6} if explicit else {} + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + result = await certify_upstream_url( + "https://mock.example/v1", + model_id="test-model", + provider_fee=fee, + sats_usd_price=0.001, + check_cache=False, + client=client, + **kwargs, + ) + row = next(row for row in result["rows"] if row["id"] == "cost.prompt_completion") + assert row["status"] == "ok", row + assert row["evidence"]["actual_total_msats"] == expected + + +class TrickleBody(httpx.AsyncByteStream): + def __init__(self) -> None: + self.closed = False + self.started = asyncio.Event() + + async def __aiter__(self) -> AsyncIterator[bytes]: + self.started.set() + for chunk in (b'{"data":', b"[]", b"}"): + await asyncio.sleep(0.03) + yield chunk + + async def aclose(self) -> None: + self.closed = True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["models", "chat", "cache"]) +async def test_probe_elapsed_deadline_closes_trickling_response(mode: str) -> None: + body = TrickleBody() + + def handle(request: httpx.Request) -> httpx.Response: + if mode == "chat" and request.method == "GET": + return httpx.Response(200, json={"data": []}) + return httpx.Response(200, stream=body) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + if mode == "cache": + result = await probe_cache( + "https://mock.example/v1", "", "test-model", client=client, timeout=0.05 + ) + assert result.statuses == [None] + assert "TimeoutError" in (result.errors[0] or "") + else: + probe = await probe_upstream( + "https://mock.example/v1", + "", + "test-model" if mode == "chat" else "", + client=client, + timeout=0.05, + ) + if mode == "chat": + assert probe.chat_status is None + assert "TimeoutError" in (probe.chat_error or "") + else: + assert probe.models_status is None + assert "TimeoutError" in (probe.models_error or "") + assert body.closed + assert not client.is_closed + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cache", [False, True]) +@pytest.mark.parametrize("owns_client", [False, True]) +async def test_probe_cancellation_closes_body_and_owned_client( + monkeypatch: pytest.MonkeyPatch, cache: bool, owns_client: bool +) -> None: + body = TrickleBody() + client = httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: httpx.Response(200, stream=body)) + ) + if owns_client: + monkeypatch.setattr(httpx, "AsyncClient", lambda **_: client) + probe = probe_cache if cache else probe_upstream + task = asyncio.create_task( + probe("https://mock.example/v1", "", "", client=None if owns_client else client) + ) + await body.started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert body.closed + assert client.is_closed == owns_client + await client.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [401, 403, 429, 500, 503]) +async def test_failed_initial_completion_skips_cache(status: int) -> None: + calls = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "test-model"}]}) + calls.append(request) + return httpx.Response(status, json={"error": "failed"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + rows = await run_live_checks( + "https://mock.example/v1", + "", + _model_from_usd_pricing("test-model", 1e-6, 2e-6, 0.001), + provider_fee=1, + sats_to_usd=0.001, + client=client, + ) + assert len(calls) == 1 + assert ( + next(row for row in rows if row["id"] == "cache.reported")["status"] == "warn" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [401, 403, 429, 500, 503]) +async def test_cache_does_not_retry_non_format_errors(status: int) -> None: + calls = [] + + def handle(request: httpx.Request) -> httpx.Response: + calls.append(request) + return httpx.Response(status, json={"error": "failed"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + result = await probe_cache( + "https://mock.example/v1", "", "test-model", client=client + ) + assert len(calls) == 1 + assert result.statuses == [status] + + +@pytest.mark.parametrize( + "usage", + [ + {"cache_creation_input_tokens": 3000}, + {"prompt_tokens_details": {"cache_creation_tokens": 3000}}, + {"prompt_tokens_details": {"cache_write_tokens": 3000}}, + {"input_tokens_details": {"cache_write_tokens": 3000}}, + ], +) +@pytest.mark.parametrize("unknown", [False, True]) +def test_known_cache_writes_are_no_hit_not_unrecognized( + usage: dict[str, Any], unknown: bool +) -> None: + usage = {"input_tokens": 10, "output_tokens": 1, **usage} + if unknown: + usage["unknown_cached_read_tokens"] = 5 + payload = {"usage": usage} + row = cache_reported_row( + CacheProbeResult( + chat_url="mock", statuses=[200, 200], payloads=[payload, payload] + ) + ) + assert row["status"] == ("fail" if unknown else "warn") + assert row["evidence"]["second_usage"]["cache_write_tokens"] == 3000 + assert row["evidence"].get("unrecognised_cache_fields", []) == ( + ["unknown_cached_read_tokens"] if unknown else [] + ) diff --git a/ui/app/providers/certification/page.tsx b/ui/app/providers/certification/page.tsx index ef1b0ef4..93274331 100644 --- a/ui/app/providers/certification/page.tsx +++ b/ui/app/providers/certification/page.tsx @@ -1,6 +1,6 @@ 'use client'; -import { useEffect, useMemo, useState } from 'react'; +import { useEffect, useMemo, useRef, useState } from 'react'; import Link from 'next/link'; import { useQueries, useQuery } from '@tanstack/react-query'; import { @@ -58,6 +58,7 @@ import { countCertificationTargets, emptyCertificationSetup, getModelsNeedingPath, + getSelectedCertificationResults, } from '@/lib/provider-certification'; import type { CertificationProgress, @@ -86,6 +87,14 @@ export default function MultiProviderCertificationPage() { Record >({}); const [runningProviderIds, setRunningProviderIds] = useState([]); + const activeRuns = useRef(new Set()); + const generation = useRef(0); + + useEffect(() => { + return () => { + generation.current += 1; + }; + }, []); const providersQuery = useQuery({ queryKey: ['upstream-providers'], @@ -180,6 +189,10 @@ export default function MultiProviderCertificationPage() { }; const runOneProvider = async (providerId: number) => { + if (activeRuns.current.has(providerId)) return; + activeRuns.current.add(providerId); + const runGeneration = generation.current; + const shouldContinue = () => generation.current === runGeneration; const setup = setups[providerId] ?? emptyCertificationSetup(); const modelRuns = providerRuns(providerId); setResultsByProvider((current) => ({ ...current, [providerId]: [] })); @@ -192,25 +205,35 @@ export default function MultiProviderCertificationPage() { providerId, modelRuns, includeCache: setup.checkCache, - onProgress: (progress) => - setProgressByProvider((current) => ({ - ...current, - [providerId]: progress, - })), - onResults: (results) => - setResultsByProvider((current) => ({ - ...current, - [providerId]: results, - })), + shouldContinue, + onProgress: (progress) => { + if (shouldContinue()) { + setProgressByProvider((current) => ({ + ...current, + [providerId]: progress, + })); + } + }, + onResults: (results) => { + if (shouldContinue()) { + setResultsByProvider((current) => ({ + ...current, + [providerId]: results, + })); + } + }, }); } finally { - setProgressByProvider((current) => ({ - ...current, - [providerId]: null, - })); - setRunningProviderIds((current) => - current.filter((id) => id !== providerId) - ); + activeRuns.current.delete(providerId); + if (shouldContinue()) { + setProgressByProvider((current) => ({ + ...current, + [providerId]: null, + })); + setRunningProviderIds((current) => + current.filter((id) => id !== providerId) + ); + } } }; @@ -241,7 +264,10 @@ export default function MultiProviderCertificationPage() { ); const allReady = selectedProviderIds.length > 0 && incompleteProviders.length === 0; - const allResults = Object.values(resultsByProvider).flat(); + const allResults = getSelectedCertificationResults( + selectedProviderIds, + resultsByProvider + ); const aggregateSummary = summarizeCertificationResults(allResults); const pendingRoutes = Math.max(totalRoutes - allResults.length, 0); diff --git a/ui/components/provider-certification-dialog.tsx b/ui/components/provider-certification-dialog.tsx index 381cda4e..ee8dd24b 100644 --- a/ui/components/provider-certification-dialog.tsx +++ b/ui/components/provider-certification-dialog.tsx @@ -80,7 +80,13 @@ export function ProviderCertificationDialog({ }; return ( - + { + if (!nextOpen) reset(); + onOpenChange(nextOpen); + }} + > Certify upstream models @@ -147,7 +153,9 @@ export function ProviderCertificationDialog({ )} {isPending - ? `Running ${progress?.modelIndex ?? 1} of ${progress?.modelTotal ?? setup.selectedModelIds.length}` + ? progress + ? `Running ${progress.modelIndex} of ${progress.modelTotal}` + : 'Finishing in-flight probe' : results.length > 0 ? `Run ${targetCount} route${targetCount === 1 ? '' : 's'} again` : `Certify ${targetCount || ''} route${targetCount === 1 ? '' : 's'}`} diff --git a/ui/hooks/use-provider-certification-runner.test.mjs b/ui/hooks/use-provider-certification-runner.test.mjs new file mode 100644 index 00000000..f2cfbeb1 --- /dev/null +++ b/ui/hooks/use-provider-certification-runner.test.mjs @@ -0,0 +1,328 @@ +import assert from 'node:assert/strict'; +import { readFileSync } from 'node:fs'; +import { test } from 'node:test'; +import { fileURLToPath } from 'node:url'; +import vm from 'node:vm'; +import ts from 'typescript'; + +function loadSource(path, imports) { + const source = readFileSync(new URL(path, import.meta.url), 'utf8'); + const { outputText } = ts.transpileModule(source, { + compilerOptions: { + module: ts.ModuleKind.CommonJS, + target: ts.ScriptTarget.ES2020, + jsx: ts.JsxEmit.ReactJSX, + }, + fileName: fileURLToPath(new URL(path, import.meta.url)), + }); + const sourceModule = { exports: {} }; + vm.runInNewContext(outputText, { + module: sourceModule, + exports: sourceModule.exports, + require(name) { + assert.ok(name in imports, `Unexpected import ${name}`); + return imports[name]; + }, + }); + return sourceModule.exports; +} + +function runnerHarness() { + const calls = []; + const pending = []; + const slots = []; + const cleanups = []; + let cursor = 0; + const react = { + useState(initial) { + const index = cursor++; + if (!(index in slots)) slots[index] = initial; + return [ + slots[index], + (next) => { + slots[index] = next; + }, + ]; + }, + useRef(initial) { + const index = cursor++; + if (!(index in slots)) slots[index] = { current: initial }; + return slots[index]; + }, + useEffect(effect) { + const index = cursor++; + if (!(index in slots)) { + slots[index] = true; + cleanups.push(effect()); + } + }, + useCallback: (fn) => fn, + }; + const exports = loadSource('./use-provider-certification-runner.ts', { + react, + '@/lib/api/services/admin': { + AdminService: { + certifyProvider(providerId, options) { + calls.push({ providerId, ...options }); + return new Promise((resolve, reject) => + pending.push({ resolve, reject }) + ); + }, + }, + }, + '@/lib/provider-certification': { + getErrorMessage: (error) => error.message, + }, + }); + return { + ...exports, + calls, + pending, + render() { + cursor = 0; + return exports.useProviderCertificationRunner(1); + }, + unmount() { + cleanups.forEach((cleanup) => cleanup?.()); + }, + }; +} + +const model = (modelId, paths = ['default']) => ({ + modelId, + targets: paths.map((path) => ({ path, label: path })), +}); +const report = { rows: [] }; +const flush = () => new Promise((resolve) => setImmediate(resolve)); + +test('reset cancels queued paths/models and blocks restart until in-flight work finishes', async () => { + const harness = runnerHarness(); + const hook = harness.render(); + const old = hook.run([model('first', ['a', 'b']), model('second')], false); + assert.equal(harness.render().isPending, true); + hook.reset(); + assert.equal(harness.render().isPending, true); + await hook.run([model('restart')], false); + assert.equal(harness.calls.length, 1); + harness.pending[0].resolve(report); + await old; + const finished = harness.render(); + assert.equal(finished.isPending, false); + assert.equal(finished.progress, null); + assert.equal(finished.results.length, 0); + assert.deepEqual( + harness.calls.map((call) => call.model_id), + ['first'] + ); + const fresh = finished.run([model('restart')], false); + harness.pending[1].resolve(report); + await fresh; + assert.deepEqual( + harness.calls.map((call) => call.model_id), + ['first', 'restart'] + ); +}); + +test('unmount cancels queued requests and suppresses stale result updates', async () => { + const harness = runnerHarness(); + const old = harness.render().run([model('first'), model('second')], false); + harness.unmount(); + harness.pending[0].resolve(report); + await old; + assert.equal(harness.calls.length, 1); + assert.equal(harness.render().results.length, 0); +}); + +test('cancellation before dispatch makes no request or progress callback', async () => { + const harness = runnerHarness(); + const progress = []; + const result = await harness.runProviderCertification({ + providerId: 1, + modelRuns: [model('first')], + includeCache: false, + shouldContinue: () => false, + onProgress: (next) => progress.push(next), + }); + assert.equal(result.length, 0); + assert.equal(harness.calls.length, 0); + assert.equal(progress.length, 0); +}); + +test('per-route failures remain isolated and normal runs retain completed results', async () => { + const harness = runnerHarness(); + const done = harness.render().run([model('first', ['a', 'b'])], true); + harness.pending[0].reject(new Error('route failed')); + await flush(); + harness.pending[1].resolve(report); + await done; + const finished = harness.render(); + assert.equal(finished.isPending, false); + assert.equal(finished.results.length, 2); + assert.equal(finished.results[0].error, 'route failed'); + assert.equal(finished.results[1].report, report); + assert.equal(harness.calls[1].check_cache, true); +}); + +test('dialog close cancels synchronously before notifying its owner', () => { + const events = []; + const jsx = (type, props) => ({ type, props }); + const components = new Proxy({}, { get: (_, name) => name }); + const imports = { + react: { + useState: (initial) => [ + typeof initial === 'function' ? initial() : initial, + () => {}, + ], + useEffect: (effect) => effect(), + }, + 'react/jsx-runtime': { jsx, jsxs: jsx }, + '@tanstack/react-query': { useQuery: () => ({}) }, + 'lucide-react': components, + '@/components/ui/dialog': components, + '@/components/ui/button': components, + '@/components/ui/badge': components, + '@/components/ui/tabs': components, + '@/components/provider-certification-results': components, + '@/components/provider-certification-setup': { + ...components, + ProviderCertificationSetupPanel: 'SetupPanel', + getCertificationModelNames: () => ({}), + }, + '@/hooks/use-provider-certification-runner': { + useProviderCertificationRunner: () => ({ + results: [], + progress: null, + isPending: false, + run: () => {}, + reset: () => events.push('reset'), + }), + }, + '@/lib/api/services/admin': { AdminService: {} }, + '@/lib/provider-certification': { + emptyCertificationSetup: () => ({ + selectedModelIds: [], + checkCache: false, + }), + buildModelRuns: () => [], + countCertificationTargets: () => 0, + getModelsNeedingPath: () => [], + }, + }; + const { ProviderCertificationDialog } = loadSource( + '../components/provider-certification-dialog.tsx', + imports + ); + const dialog = ProviderCertificationDialog({ + open: true, + provider: { id: 1, provider_type: 'generic' }, + onOpenChange: () => events.push('owner'), + }); + dialog.props.onOpenChange(false); + assert.deepEqual(events, ['reset', 'owner']); +}); + +test('multi-provider page guards duplicate starts and cancels queued work on unmount', async () => { + const harness = runnerHarness(); + const cleanups = []; + const updates = []; + let stateIndex = 0; + const setup = { + selectedModelIds: ['first', 'second'], + pathModes: {}, + selectedModelPaths: {}, + checkCache: false, + }; + const initialStates = [[1], 1, { 1: setup }]; + const react = { + useState(initial) { + const index = stateIndex++; + return [ + index < 3 ? initialStates[index] : initial, + (next) => updates.push(next), + ]; + }, + useMemo: (fn) => fn(), + useRef: (initial) => ({ current: initial }), + useEffect: (effect) => cleanups.push(effect()), + }; + const helpers = loadSource('../lib/provider-certification.ts', { + '@/lib/api/errors': { getApiErrorMessage: () => '' }, + }); + const jsx = (type, props) => ({ type, props }); + const components = new Proxy({}, { get: (_, name) => name }); + const imports = { + react, + 'react/jsx-runtime': { jsx, jsxs: jsx }, + 'next/link': { default: 'Link' }, + '@tanstack/react-query': { + useQuery: () => ({ data: [{ id: 1, provider_type: 'generic' }] }), + useQueries: () => [{ data: { certification_paths: {} } }], + }, + 'lucide-react': components, + '@/components/provider-certification-results': { + summarizeCertificationResults: () => ({}), + ProviderCertificationResults: 'Results', + }, + '@/components/provider-certification-setup': { + getCertificationModelNames: () => ({}), + ProviderCertificationSetupPanel: 'Setup', + }, + '@/hooks/use-provider-certification-runner': harness, + '@/lib/api/services/admin': { AdminService: {} }, + '@/lib/provider-certification': helpers, + '@/lib/utils': { cn: () => '' }, + }; + for (const name of [ + 'app-page-shell', + 'page-header', + 'ui/badge', + 'ui/button', + 'ui/card', + 'ui/checkbox', + 'ui/command', + 'ui/popover', + 'ui/select', + 'ui/tabs', + ]) + imports[`@/components/${name}`] = components; + const { default: Page } = loadSource( + '../app/providers/certification/page.tsx', + imports + ); + const nodes = []; + const visit = (node) => { + if (!node || typeof node !== 'object') return; + if (Array.isArray(node)) return node.forEach(visit); + nodes.push(node); + visit(node.props?.children); + }; + visit(Page()); + const runAll = nodes.find( + (node) => node.props?.onClick?.name === 'runAllProviders' + ); + assert.ok(runAll); + runAll.props.onClick(); + runAll.props.onClick(); + assert.equal(harness.calls.length, 1); + cleanups.forEach((cleanup) => cleanup?.()); + const beforeCompletion = updates.length; + harness.pending[0].resolve(report); + await flush(); + assert.equal(harness.calls.length, 1); + assert.equal(updates.length, beforeCompletion); +}); + +test('selected-provider aggregate excludes deselected providers and restores them on reselection', () => { + const { getSelectedCertificationResults } = loadSource( + '../lib/provider-certification.ts', + { + '@/lib/api/errors': { getApiErrorMessage: () => '' }, + } + ); + const a = { providerId: 1, resultKey: 'a' }; + const results = { 1: [a] }; + const selected = getSelectedCertificationResults([2], results); + assert.equal(selected.length, 0); + assert.equal(Math.max(1 - selected.length, 0), 1); + assert.deepEqual([...getSelectedCertificationResults([1, 2], results)], [a]); +}); diff --git a/ui/hooks/use-provider-certification-runner.ts b/ui/hooks/use-provider-certification-runner.ts index fb7a9f7d..8bd117aa 100644 --- a/ui/hooks/use-provider-certification-runner.ts +++ b/ui/hooks/use-provider-certification-runner.ts @@ -14,6 +14,7 @@ interface RunProviderCertificationOptions { providerId: number; modelRuns: ModelRun[]; includeCache: boolean; + shouldContinue?: () => boolean; onProgress?: (progress: CertificationProgress | null) => void; onResults?: (results: ModelCertificationResult[]) => void; } @@ -22,12 +23,14 @@ export async function runProviderCertification({ providerId, modelRuns, includeCache, + shouldContinue = () => true, onProgress, onResults, }: RunProviderCertificationOptions): Promise { const completed: ModelCertificationResult[] = []; for (const [index, run] of modelRuns.entries()) { + if (!shouldContinue()) break; onProgress?.({ modelId: run.modelId, modelIndex: index + 1, @@ -38,6 +41,7 @@ export async function runProviderCertification({ // parallel paths multiply that spend and the admin request load. const batch: ModelCertificationResult[] = []; for (const target of run.targets) { + if (!shouldContinue()) break; const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`; try { const report = await AdminService.certifyProvider(providerId, { @@ -62,11 +66,12 @@ export async function runProviderCertification({ }); } } + if (!shouldContinue()) break; completed.push(...batch); onResults?.([...completed]); } - onProgress?.(null); + if (shouldContinue()) onProgress?.(null); return completed; } @@ -76,6 +81,7 @@ export function useProviderCertificationRunner(providerId: number) { const [isPending, setIsPending] = useState(false); const generation = useRef(0); const mounted = useRef(true); + const active = useRef(false); useEffect(() => { mounted.current = true; @@ -89,11 +95,14 @@ export function useProviderCertificationRunner(providerId: number) { generation.current += 1; setResults([]); setProgress(null); - setIsPending(false); + // Already dispatched probes may still spend credit; keep them pending. + setIsPending(active.current); }, []); const run = useCallback( async (modelRuns: ModelRun[], includeCache: boolean) => { + if (active.current || !mounted.current) return []; + active.current = true; const runGeneration = generation.current + 1; generation.current = runGeneration; setResults([]); @@ -104,6 +113,8 @@ export function useProviderCertificationRunner(providerId: number) { providerId, modelRuns, includeCache, + shouldContinue: () => + mounted.current && generation.current === runGeneration, onProgress: (nextProgress) => { if (mounted.current && generation.current === runGeneration) { setProgress(nextProgress); @@ -116,7 +127,8 @@ export function useProviderCertificationRunner(providerId: number) { }, }); } finally { - if (mounted.current && generation.current === runGeneration) { + active.current = false; + if (mounted.current) { setProgress(null); setIsPending(false); } diff --git a/ui/lib/provider-certification.ts b/ui/lib/provider-certification.ts index 1b9a0d20..dff4478c 100644 --- a/ui/lib/provider-certification.ts +++ b/ui/lib/provider-certification.ts @@ -104,6 +104,12 @@ export const buildModelRuns = ( export const countCertificationTargets = (runs: ModelRun[]): number => runs.reduce((total, run) => total + run.targets.length, 0); +export const getSelectedCertificationResults = ( + providerIds: number[], + resultsByProvider: Record +): ModelCertificationResult[] => + providerIds.flatMap((providerId) => resultsByProvider[providerId] ?? []); + export const getCertificationResultStatus = ( result: ModelCertificationResult ): CertificationStatus | 'error' => { From ffcabba06099bf64c07e4146ef55c653fc3f39b6 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 2 Oct 2026 17:07:03 +0200 Subject: [PATCH 14/18] fix: shape certification probes through provider request hooks --- routstr/core/admin.py | 54 +++-- routstr/upstream/certification.py | 128 +++++++++--- routstr/upstream/certification_cache.py | 61 +++--- tests/integration/test_certify_alias_paths.py | 3 +- tests/integration/test_certify_endpoint.py | 9 +- .../test_certify_provider_shapes.py | 185 ++++++++++++++++++ .../test_certification_review_regressions.py | 4 +- .../provider-certification-setup.tsx | 6 + ui/lib/api/services/admin.ts | 1 - 9 files changed, 353 insertions(+), 98 deletions(-) create mode 100644 tests/integration/test_certify_provider_shapes.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index bf3aac15..31140ff9 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1523,7 +1523,6 @@ async def get_upstream_provider_report(provider_id: str) -> dict[str, object]: class CertifyRequest(BaseModel): model_id: str | None = None model_path: str | None = None - timeout_seconds: float | None = None check_cache: bool = True @@ -1538,8 +1537,9 @@ async def certify_upstream_provider( Unlike the read-only ``GET …/report``, this probes the upstream over the network and runs the node's cost engine on the real response. It never - enters the billing path, so it costs at most one completion's worth of - upstream credit and nothing from the node's wallet. + enters the billing path, so it costs nothing from the node's wallet. Its + upstream spend is a one-token completion, plus two or three one-token + completions on a ~4.4k-token prompt when ``check_cache`` is set. Returns the read-only report's four ``pricing.*`` rows (re-derived here so the certification is self-contained), the live rows from @@ -1547,7 +1547,6 @@ async def certify_upstream_provider( operator-facing goals. """ from ..upstream.certification import ( - MAX_PROBE_TIMEOUT_SECONDS, PROBE_TIMEOUT_SECONDS, build_checklist, run_live_checks, @@ -1565,6 +1564,7 @@ async def certify_upstream_provider( enabled_rows = list(result.all()) endpoint_tag: str | None = None + path_model_id: str | None = None selected_path: ModelPathRow | None = None if payload.model_path is not None: from ..upstream.model_paths import decode_model_path @@ -1585,6 +1585,7 @@ async def certify_upstream_provider( detail="Model path is not available for this provider", ) endpoint_tag = selector.endpoint_tag + path_model_id = selector.model_id evaluations = [ _evaluate_model_row(row, provider, provider_pk) for row in enabled_rows @@ -1609,6 +1610,12 @@ async def certify_upstream_provider( from ..proxy import get_candidates, get_upstreams + # The live upstream instance shapes the probes exactly like the proxy's + # own requests (paths, auth headers, query params, model-name transforms). + upstream_obj = next( + (u for u in get_upstreams() if getattr(u, "db_id", None) == provider_pk), + None, + ) model_obj = None if model_id: try: @@ -1633,20 +1640,15 @@ async def certify_upstream_provider( # The active upstream cache carries the same fee-adjusted USD and sats # pricing used by the proxy, so it is the authoritative fallback for a # pre-configuration certification probe. - if model_obj is None: - for upstream in get_upstreams(): - if getattr(upstream, "db_id", None) != provider_pk: - continue - model_obj = next( - ( - model - for model in upstream.get_cached_models() - if model.id == model_id or model.forwarded_model_id == model_id - ), - None, - ) - if model_obj is not None: - break + if model_obj is None and upstream_obj is not None: + model_obj = next( + ( + model + for model in upstream_obj.get_cached_models() + if model.id == model_id or model.forwarded_model_id == model_id + ), + None, + ) if selected_path is not None: from ..upstream.model_paths import exposed_model_id @@ -1654,7 +1656,8 @@ async def certify_upstream_provider( if ( payload.model_id is None or selected_id is None - or selected_id.lower() != selector.model_id.lower() + or path_model_id is None + or selected_id.lower() != path_model_id.lower() ): raise HTTPException( status_code=400, @@ -1726,23 +1729,18 @@ async def certify_upstream_provider( provider.provider_fee, sats_to_usd, ) - # Clamp the admin-supplied timeout per upstream call. The run makes up - # to five calls, so the request can stay open for up to five times it. - 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) + # The timeout applies per upstream call. The run makes up to five + # calls, so the request can stay open for up to five times it. live_rows = await run_live_checks( provider.base_url, provider.api_key, model_obj, provider_fee=provider.provider_fee, sats_to_usd=sats_to_usd, - timeout=timeout, + timeout=PROBE_TIMEOUT_SECONDS, check_cache=payload.check_cache, endpoint_tag=endpoint_tag, + upstream=upstream_obj, ) rows = pricing_rows + live_rows diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 1051b019..2989e287 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -4,7 +4,9 @@ Extends the read-only pricing rows, which never touch the network, with the ones that must: a ``/models`` heartbeat and a one-token completion. Probes call the upstream directly with ``httpx``, never through the node's -billing path — no reservation, no Cashu, at most one token of upstream spend. +billing path — no reservation, no Cashu. Upstream spend is one one-token +completion, plus two or three one-token completions on a ~4.4k-token prompt +when the cache checks are enabled (per certified model path). They sit behind ``POST …/certify`` rather than the read-only ``GET …/report`` because they can block for the length of the timeout. """ @@ -25,13 +27,14 @@ from urllib.parse import urlparse import httpx from ..core.logging import get_logger -from ..payment.cost_calculation import calculate_cost +from ..payment.cost_calculation import _resolve_usd_cost, calculate_cost from ..payment.rates import coerce_rate from ..payment.usage import normalize_usage from .model_paths import is_openrouter_base_url if TYPE_CHECKING: from ..payment.models import Model + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -44,9 +47,6 @@ TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"} # Bounded so a dead upstream fails the row rather than wedging the request. PROBE_TIMEOUT_SECONDS = 15.0 -# Ceiling for the caller-supplied timeout override. -MAX_PROBE_TIMEOUT_SECONDS = 60.0 - # The cheapest request that still exercises the usage/cost path. PROBE_MAX_TOKENS = 1 PROBE_PROMPT = "ping" @@ -181,6 +181,72 @@ class ProbeResult: chat_latency_ms: float | None = None +@dataclass +class ProbeShape: + """Where the probe calls go and how they are authenticated.""" + + models_url: str + chat_url: str + headers: dict[str, str] + models_params: dict[str, str] + chat_params: dict[str, str] + + +def probe_shape( + base_url: str, + api_key: str, + upstream: "BaseUpstreamProvider | None" = None, + model: "Model | None" = None, +) -> ProbeShape: + """The URLs, headers and query params a probe sends. + + With the node's upstream instance, use the hooks + ``BaseUpstreamProvider.forward_request`` uses, so the probe reaches what + the proxy reaches (Azure's deployment path and ``api-key``, Gemini's + ``/openai`` base, Ollama's ``/v1``). Without one (the CLI), assume a plain + OpenAI-compatible base URL. + """ + if upstream is None: + base = base_url.rstrip("/") + headers = {"Content-Type": "application/json"} + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + return ProbeShape(f"{base}/models", f"{base}/chat/completions", headers, {}, {}) + chat_path = upstream.normalize_request_path("v1/chat/completions", model) + # Azure lists models under ``/openai/models``, not at its endpoint root. + if upstream.provider_type == "azure": + models_path = "openai/models" + else: + models_path = upstream.normalize_request_path("v1/models") + return ProbeShape( + models_url=upstream.build_request_url(models_path), + chat_url=upstream.build_request_url(chat_path, model), + headers=upstream.prepare_headers({"content-type": "application/json"}), + models_params=dict(upstream.prepare_params(models_path, None)), + chat_params=dict(upstream.prepare_params(chat_path, None)), + ) + + +def shape_body( + body: dict[str, Any], + upstream: "BaseUpstreamProvider | None" = None, + model: "Model | None" = None, +) -> Any: + """The JSON body the proxy would forward, model-name transforms included. + + ``prepare_request_body`` rewrites ``model`` from ``model.id``; the probe + keeps the id it chose (``forwarded_model_id`` first) and only applies the + provider's own name transform to it. + """ + if upstream is None or model is None: + return body + shaped = upstream.prepare_request_body(json.dumps(body).encode(), model) + data = json.loads(shaped) if shaped else dict(body) + if isinstance(data, dict) and isinstance(body.get("model"), str): + data["model"] = upstream.transform_model_name(body["model"]) + return data + + async def probe_upstream( base_url: str, api_key: str, @@ -189,22 +255,22 @@ async def probe_upstream( endpoint_tag: str | None = None, client: httpx.AsyncClient | None = None, timeout: float = PROBE_TIMEOUT_SECONDS, + upstream: "BaseUpstreamProvider | None" = None, + model: "Model | None" = None, ) -> ProbeResult: """Call the upstream's ``/models`` and a one-token completion. Each HTTP call, including its body read, has an elapsed-time deadline. A transport failure is a ``fail`` row, not a failed admin request. """ - base = base_url.rstrip("/") + shape = probe_shape(base_url, api_key, upstream, model) result = ProbeResult( base_url=base_url, - models_url=f"{base}/models", - chat_url=f"{base}/chat/completions", + models_url=shape.models_url, + chat_url=shape.chat_url, endpoint_tag=endpoint_tag, ) - headers = {"Content-Type": "application/json"} - if api_key: - headers["Authorization"] = f"Bearer {api_key}" + headers = shape.headers owns_client = client is None if client is None: @@ -214,7 +280,9 @@ async def probe_upstream( started = time.monotonic() try: async with asyncio.timeout(timeout): - response = await client.get(result.models_url, headers=headers) + response = await client.get( + result.models_url, headers=headers, params=shape.models_params + ) result.models_status = response.status_code result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) try: @@ -249,7 +317,10 @@ async def probe_upstream( try: async with asyncio.timeout(timeout): response = await client.post( - result.chat_url, json=request_body, headers=headers + result.chat_url, + json=shape_body(request_body, upstream, model), + headers=headers, + params=shape.chat_params, ) result.chat_status = response.status_code result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) @@ -497,27 +568,14 @@ def _truncate(value: Any, limit: int = 400) -> Any: def _reported_usd_cost(payload: dict[str, Any]) -> float: """The upstream-reported USD cost, or 0.0 when it reported none. - Mirrors ``_resolve_usd_cost``'s priority and shares ``coerce_rate``, so - this helper and the engine agree on *whether* a cost was reported; only - the arithmetic below is re-derived independently. + Uses the engine's own ``_resolve_usd_cost`` so both agree on *which* + figure is the cost (PPQ.AI BYOK bills ``upstream_inference_cost`` plus the + fee); only the arithmetic below is re-derived independently. """ usage = payload.get("usage") if not isinstance(usage, dict): return 0.0 - cost_details = usage.get("cost_details") - if isinstance(cost_details, dict): - total = coerce_rate(cost_details.get("total_cost")) - if total is not None and total > 0: - return total - inference = coerce_rate(cost_details.get("upstream_inference_cost")) - if inference is not None and inference > 0 and usage.get("is_byok"): - return inference + (coerce_rate(usage.get("cost")) or 0.0) - for source in (usage, payload): - for field in ("total_cost", "cost"): - value = coerce_rate(source.get(field)) - if value is not None and value > 0: - return value - return 0.0 + return _resolve_usd_cost(usage, payload) def _fixed_token_pricing_active() -> bool: @@ -758,11 +816,14 @@ async def run_live_checks( pricing_known: bool = True, check_cache: bool = True, endpoint_tag: str | None = None, + upstream: "BaseUpstreamProvider | None" = None, ) -> list[dict[str, Any]]: """Probe one upstream and build the live/derived rows. - ``check_cache`` adds the prompt-cache and margin rows, which cost two - more completions against a long prompt. + ``check_cache`` adds the prompt-cache and margin rows, which cost two or + three more completions against a long prompt. ``upstream`` shapes the + probes like the proxy's own requests; without it they assume a plain + OpenAI-compatible base URL. """ probe = await probe_upstream( base_url, @@ -771,6 +832,8 @@ async def run_live_checks( endpoint_tag=endpoint_tag, client=client, timeout=timeout, + upstream=upstream, + model=model, ) rows = [ safe_row( @@ -863,6 +926,7 @@ async def run_live_checks( timeout=timeout, pricing_known=pricing_known, endpoint_tag=endpoint_tag, + upstream=upstream, ) ) return rows diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index d8ed2dad..9c267b69 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -37,11 +37,14 @@ from .certification import ( _reported_usd_cost, _token_rates, certification_row, + probe_shape, safe_row, + shape_body, ) if TYPE_CHECKING: from ..payment.models import Model + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -129,14 +132,15 @@ def _request_body( async def _post_completion( client: httpx.AsyncClient, url: str, - body: dict[str, Any], + body: Any, headers: dict[str, str], + params: dict[str, str], timeout: float, ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: started = time.monotonic() try: async with asyncio.timeout(timeout): - response = await client.post(url, json=body, headers=headers) + response = await client.post(url, json=body, headers=headers, params=params) except Exception as exc: # noqa: BLE001 - transport failure is a row status latency = round((time.monotonic() - started) * 1000, 2) return None, None, f"{type(exc).__name__}: {exc}", latency @@ -178,6 +182,8 @@ async def probe_cache( endpoint_tag: str | None = None, client: httpx.AsyncClient | None = None, timeout: float = PROBE_TIMEOUT_SECONDS, + upstream: "BaseUpstreamProvider | None" = None, + model: "Model | None" = None, ) -> CacheProbeResult: """Send the same long prompt twice. @@ -186,49 +192,39 @@ async def probe_cache( retry on HTTP 400/422, and the second call mirrors whichever format succeeded. Each call's elapsed deadline includes the response body. """ - base = base_url.rstrip("/") - result = CacheProbeResult( - chat_url=f"{base}/chat/completions", endpoint_tag=endpoint_tag - ) - headers = {"Content-Type": "application/json"} - if api_key: - headers["Authorization"] = f"Bearer {api_key}" + shape = probe_shape(base_url, api_key, upstream, model) + result = CacheProbeResult(chat_url=shape.chat_url, endpoint_tag=endpoint_tag) prefix = cache_probe_prefix() owns_client = client is None - if client is None: - client = httpx.AsyncClient(timeout=timeout) - try: - first = await _post_completion( - client, + http = client if client is not None else httpx.AsyncClient(timeout=timeout) + + async def post( + fmt: str, + ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: + body = _request_body(model_id, prefix, fmt, endpoint_tag) + return await _post_completion( + http, result.chat_url, - _request_body(model_id, prefix, "cache_control", endpoint_tag), - headers, + shape_body(body, upstream, model), + shape.headers, + shape.chat_params, timeout, ) + + try: + first = await post("cache_control") if first[0] in (400, 422): result.request_format = "plain" - first = await _post_completion( - client, - result.chat_url, - _request_body(model_id, prefix, "plain", endpoint_tag), - headers, - timeout, - ) + first = await post("plain") _record(result, first) if not _is_2xx(first[0]): return result - second = await _post_completion( - client, - result.chat_url, - _request_body(model_id, prefix, result.request_format, endpoint_tag), - headers, - timeout, - ) + second = await post(result.request_format) _record(result, second) finally: if owns_client: - await client.aclose() + await http.aclose() return result @@ -564,6 +560,7 @@ async def run_cache_checks( timeout: float = PROBE_TIMEOUT_SECONDS, pricing_known: bool = True, endpoint_tag: str | None = None, + upstream: "BaseUpstreamProvider | None" = None, ) -> list[dict[str, Any]]: """Run the cache probe and build the three cache/margin rows.""" probe = await probe_cache( @@ -573,6 +570,8 @@ async def run_cache_checks( endpoint_tag=endpoint_tag, client=client, timeout=timeout, + upstream=upstream, + model=model, ) cost_data = await _price_payload(probe.second_payload, model, provider_fee) return [ diff --git a/tests/integration/test_certify_alias_paths.py b/tests/integration/test_certify_alias_paths.py index 44e10f7c..6786a8a6 100644 --- a/tests/integration/test_certify_alias_paths.py +++ b/tests/integration/test_certify_alias_paths.py @@ -12,7 +12,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.db import ModelPathRow from routstr.proxy import reinitialize_upstreams from routstr.upstream.model_paths import encode_model_path -from tests.integration.test_certify_endpoint import ( + +from .test_certify_endpoint import ( _admin_headers, _make_provider, _model_row, diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index ac38a9e0..a57ce706 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -22,6 +22,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.admin import admin_sessions from routstr.core.db import ModelPathRow, ModelRow, UpstreamProviderRow from routstr.proxy import reinitialize_upstreams +from routstr.upstream.generic import GenericUpstreamProvider from routstr.upstream.model_paths import encode_model_path @@ -805,9 +806,7 @@ async def test_certify_explicit_discovered_model_without_override( 0.0005, ) - class FakeUpstream: - db_id = provider.id - + class FakeUpstream(GenericUpstreamProvider): def get_cached_models(self) -> list[Model]: return [remote_model] @@ -820,7 +819,9 @@ async def test_certify_explicit_discovered_model_without_override( return_value=Response(200, json=_mock_chat_response(model="remote-model")) ) - with patch("routstr.proxy.get_upstreams", return_value=[FakeUpstream()]): + fake = FakeUpstream(base_url=provider.base_url, api_key=provider.api_key) + fake.db_id = provider.id + with patch("routstr.proxy.get_upstreams", return_value=[fake]): resp = await integration_client.post( f"/admin/api/upstream-providers/{provider.id}/certify", headers=_admin_headers(), diff --git a/tests/integration/test_certify_provider_shapes.py b/tests/integration/test_certify_provider_shapes.py new file mode 100644 index 00000000..ee7ab153 --- /dev/null +++ b/tests/integration/test_certify_provider_shapes.py @@ -0,0 +1,185 @@ +"""Certification probes must send the request the proxy would send. + +Each provider type reshapes requests through its hooks (paths, auth headers, +query params, model-name transforms). The mocked upstream here answers only +the proxy-shaped request, so a probe that hand-builds an OpenAI-style call +fails ``endpoint.reachable`` / ``usage.capture`` instead of passing. +""" + +from __future__ import annotations + +import json +from dataclasses import dataclass, field +from unittest.mock import patch + +import pytest +import respx +from httpx import AsyncClient, Response +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + +from .test_certify_endpoint import ( + _admin_headers, + _find_row, + _mock_chat_response, + _model_row, + _pin_sats_usd, # noqa: F401 - autouse: pins the sats/USD quote +) + + +@dataclass(frozen=True) +class Shape: + provider_type: str + base_url: str + model_id: str + models_url: str + chat_url: str + upstream_model: str + auth_header: tuple[str, str] + params: dict[str, str] = field(default_factory=dict) + api_version: str | None = None + + +SHAPES = [ + Shape( + provider_type="azure", + base_url="https://res.openai.azure.com", + model_id="gpt-4o", + models_url="https://res.openai.azure.com/openai/models", + chat_url=( + "https://res.openai.azure.com/openai/deployments/gpt-4o/chat/completions" + ), + upstream_model="gpt-4o", + auth_header=("api-key", "test-key"), + params={"api-version": "2024-10-21"}, + api_version="2024-10-21", + ), + Shape( + provider_type="gemini", + base_url="https://generativelanguage.googleapis.com/v1beta", + model_id="gemini-2.5-flash", + models_url="https://generativelanguage.googleapis.com/v1beta/openai/models", + chat_url=( + "https://generativelanguage.googleapis.com/v1beta/openai/chat/completions" + ), + upstream_model="gemini-2.5-flash", + auth_header=("authorization", "Bearer test-key"), + ), + Shape( + provider_type="ollama", + base_url="http://ollama.test:11434", + model_id="llama3", + models_url="http://ollama.test:11434/v1/models", + chat_url="http://ollama.test:11434/v1/chat/completions", + upstream_model="llama3", + auth_header=("authorization", "Bearer test-key"), + ), + Shape( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + model_id="claude-sonnet-4.5", + models_url="https://api.anthropic.com/v1/models", + chat_url="https://api.anthropic.com/v1/chat/completions", + upstream_model="claude-sonnet-4-5-20250929", + auth_header=("authorization", "Bearer test-key"), + ), +] + + +async def _seed(session: AsyncSession, shape: Shape) -> int: + provider = UpstreamProviderRow( + provider_type=shape.provider_type, + base_url=shape.base_url, + api_key="test-key", + api_version=shape.api_version, + provider_fee=1.0, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + assert provider.id is not None + session.add(_model_row(provider.id, model_id=shape.model_id)) + await session.commit() + with patch("routstr.payment.models.sats_usd_price", return_value=0.0005): + await reinitialize_upstreams() + return provider.id + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("shape", SHAPES, ids=[s.provider_type for s in SHAPES]) +async def test_certify_matches_proxy_request( + shape: Shape, + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + with respx.mock(assert_all_called=False) as mock: + provider_id = await _seed(integration_session, shape) + models_route = mock.get(shape.models_url).mock( + return_value=Response(200, json={"data": [{"id": shape.model_id}]}) + ) + chat_route = mock.post(shape.chat_url).mock( + return_value=Response(200, json=_mock_chat_response(model=shape.model_id)) + ) + + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={"model_id": shape.model_id, "check_cache": True}, + ) + + assert resp.status_code == 200, resp.text + rows = resp.json()["rows"] + assert _find_row(rows, "endpoint.reachable")["status"] == "ok" + assert _find_row(rows, "usage.capture")["status"] == "ok" + assert _find_row(rows, "cost.prompt_completion")["status"] == "ok" + + assert models_route.call_count == 1 + # The one-token probe plus both cache-probe completions. + assert chat_route.call_count == 3 + header, value = shape.auth_header + for call in [*models_route.calls, *chat_route.calls]: + assert call.request.headers.get(header) == value + for key, expected in shape.params.items(): + assert call.request.url.params.get(key) == expected + for call in chat_route.calls: + assert json.loads(call.request.content)["model"] == shape.upstream_model + + +def test_shape_body_keeps_a_single_cache_control_marker() -> None: + """The cache probe's own marker must not be stamped a second time.""" + from routstr.payment.models import Architecture, Model, Pricing + from routstr.upstream.anthropic import AnthropicUpstreamProvider + from routstr.upstream.certification import shape_body + from routstr.upstream.certification_cache import _request_body + + model = Model( + id="claude-sonnet-4.5", + name="claude-sonnet-4.5", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="unknown", + instruct_type=None, + ), + pricing=Pricing(prompt=1e-6, completion=2e-6), + sats_pricing=None, + per_request_limits=None, + top_provider=None, + enabled=True, + upstream_provider_id=1, + canonical_slug=None, + ) + upstream = AnthropicUpstreamProvider(api_key="test-key") + body = _request_body(model.id, "prefix", "cache_control", None) + + shaped = shape_body(body, upstream, model) + + assert json.dumps(shaped).count('"cache_control"') == 1 + assert shaped["model"] == "claude-sonnet-4-5-20250929" diff --git a/tests/unit/test_certification_review_regressions.py b/tests/unit/test_certification_review_regressions.py index 81dac641..92632b45 100644 --- a/tests/unit/test_certification_review_regressions.py +++ b/tests/unit/test_certification_review_regressions.py @@ -145,7 +145,9 @@ async def test_standalone_preserves_fee_cache_rates_and_usd_fee( return httpx.Response(200, json={"data": [{"id": "test-model"}]}) return httpx.Response(200, json={"model": "test-model", "usage": usage}) - kwargs = {"prompt_price": 1e-6, "completion_price": 2e-6} if explicit else {} + kwargs: dict[str, Any] = ( + {"prompt_price": 1e-6, "completion_price": 2e-6} if explicit else {} + ) async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: result = await certify_upstream_url( "https://mock.example/v1", diff --git a/ui/components/provider-certification-setup.tsx b/ui/components/provider-certification-setup.tsx index 2baa5415..6c2f0a98 100644 --- a/ui/components/provider-certification-setup.tsx +++ b/ui/components/provider-certification-setup.tsx @@ -304,6 +304,12 @@ export function ProviderCertificationSetupPanel({ Probe prompt caching and margin
    + {setup.checkCache && ( +

    + Sends 2–3 extra completions with a ~4.4k-token prompt per model path, + billed by the upstream. +

    + )}
    ); } diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index 0ae0e277..b3781f6e 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -79,7 +79,6 @@ export type ProviderCertification = z.infer; export type CertifyProviderRequest = { model_id?: string; model_path?: string; - timeout_seconds?: number; check_cache?: boolean; }; From a0aad288aac8a4254299c27e38af94ae562ae449 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 2 Oct 2026 20:05:39 +0200 Subject: [PATCH 15/18] fix: cover certification review leftovers and shape model tests like the proxy --- routstr/core/admin.py | 7 +- routstr/payment/models.py | 56 ++++++-- routstr/payment/usage.py | 17 +-- routstr/upstream/certification.py | 128 ++++++++++++------ routstr/upstream/certification_cache.py | 62 +++++++-- routstr/upstream/model_paths.py | 4 +- tests/integration/test_certify_endpoint.py | 18 ++- .../test_certify_provider_shapes.py | 46 +++++++ .../test_model_test_endpoint_security.py | 12 +- tests/unit/test_certification_cache.py | 41 ++++++ tests/unit/test_certification_hardening.py | 39 +----- tests/unit/test_certification_token_limit.py | 123 +++++++++++++++++ 12 files changed, 436 insertions(+), 117 deletions(-) create mode 100644 tests/unit/test_certification_token_limit.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 2027cc16..957684af 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1720,10 +1720,14 @@ async def certify_upstream_provider( status_code=503, detail="sats/USD price is not initialized yet; retry shortly", ) + # The proxy reserves and token-bills a pinned request with the model's + # own pricing, so the cost rows use it too; the path's advertised + # endpoint rates are only compared against it in the margin row. + advertised_model = None if selected_path is not None: from ..upstream.model_paths import apply_model_path_pricing - model_obj = apply_model_path_pricing( + advertised_model = apply_model_path_pricing( model_obj, selected_path, provider.provider_fee, @@ -1741,6 +1745,7 @@ async def certify_upstream_provider( check_cache=payload.check_cache, endpoint_tag=endpoint_tag, upstream=upstream_obj, + advertised_model=advertised_model, ) rows = pricing_rows + live_rows diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 8d7788a1..731c6de5 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -655,6 +655,44 @@ class ModelTestRequest(V2BaseModel): request_data: dict +def _model_test_target( + provider: UpstreamProviderRow, + model_row: ModelRow, + endpoint_path: str, + model_id: str, +) -> tuple[str, dict[str, str], dict[str, str], str]: + """URL, headers, query params and model id for a model test, shaped like + the proxy's. + + With the provider's live upstream instance, use the hooks + ``forward_request`` uses (Azure's deployment path, ``api-key`` and + ``api-version``, Gemini's ``/openai`` base, Ollama's ``/v1``, model-name + transforms). Without one, assume a plain OpenAI-compatible base URL. + """ + from ..proxy import get_upstreams + + upstream = next( + (u for u in get_upstreams() if getattr(u, "db_id", None) == provider.id), + None, + ) + if upstream is None: + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {provider.api_key}", + } + url = f"{provider.base_url.rstrip('/')}/{endpoint_path}" + return url, headers, {}, model_id + + model_obj = _build_model_from_row(model_row, False, provider.provider_fee) + path = upstream.normalize_request_path(f"v1/{endpoint_path}", model_obj) + return ( + upstream.build_request_url(path, model_obj), + upstream.prepare_headers({"content-type": "application/json"}), + dict(upstream.prepare_params(path, None)), + upstream.transform_model_name(model_id), + ) + + @models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) async def test_model( payload: ModelTestRequest, @@ -688,8 +726,10 @@ async def test_model( raise HTTPException(status_code=400, detail="Unsupported endpoint_type") actual_model_id = model_row.forwarded_model_id or model_row.id - request_data = dict(payload.request_data) - request_data["model"] = actual_model_id + url, headers, params, upstream_model_id = _model_test_target( + provider, model_row, endpoint_path, actual_model_id + ) + request_data = {**payload.request_data, "model": upstream_model_id} try: request_size = len(json.dumps(request_data).encode("utf-8")) @@ -698,9 +738,6 @@ async def test_model( if request_size > _MODEL_TEST_MAX_REQUEST_BYTES: raise HTTPException(status_code=413, detail="request_data too large") - base_url = provider.base_url.rstrip("/") - url = f"{base_url}/{endpoint_path}" - logger.info( "admin model test", extra={ @@ -712,14 +749,11 @@ async def test_model( }, ) - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {provider.api_key}", - } - try: async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post(url, json=request_data, headers=headers) + response = await client.post( + url, json=request_data, headers=headers, params=params + ) try: response_data = response.json() except Exception: diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index d69ab55c..08d675ad 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -38,8 +38,6 @@ 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 @@ -53,25 +51,18 @@ class NormalizedUsage(BaseModel): def parse_token_count(value: object) -> int: - """Parse a token count from various formats (int, float, str, bool). - - ``json.loads`` accepts bare ``Infinity``/``NaN`` and overflows ``1e999`` to - ``inf``, so an upstream can put them on the wire. ``int()`` raises on both, - which would turn a billing path into a 500; reject them like - ``is_usable_rate`` does instead. - """ + """Parse a token count from various formats (int, float, str, bool).""" if isinstance(value, bool): return 0 if isinstance(value, int): return max(0, value) if isinstance(value, float): - return max(0, int(value)) if math.isfinite(value) else 0 + return max(0, int(value)) if isinstance(value, str): try: - parsed = float(value) - except (ValueError, OverflowError): + return max(0, int(float(value))) + except ValueError: 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 2989e287..cb4d83e4 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -179,6 +179,19 @@ class ProbeResult: chat_payload: dict[str, Any] | None = None chat_error: str | None = None chat_latency_ms: float | None = None + # ``max_completion_tokens`` once the upstream rejected ``max_tokens``. + token_limit_field: str = "max_tokens" + + +def wants_max_completion_tokens(status: int | None, payload: Any) -> bool: + """Whether a 400 names ``max_completion_tokens`` as the field to use. + + OpenAI's o-series and gpt-5 reject ``max_tokens`` on chat completions with + "Unsupported parameter: 'max_tokens' ... Use 'max_completion_tokens'". + """ + if status != 400 or payload is None: + return False + return "max_completion_tokens" in json.dumps(payload, default=str) @dataclass @@ -260,8 +273,10 @@ async def probe_upstream( ) -> ProbeResult: """Call the upstream's ``/models`` and a one-token completion. - Each HTTP call, including its body read, has an elapsed-time deadline. - A transport failure is a ``fail`` row, not a failed admin request. + A completion refused with a 400 naming ``max_completion_tokens`` is + retried once with that field (OpenAI o-series, gpt-5). Each HTTP call, + including its body read, has an elapsed-time deadline. A transport + failure is a ``fail`` row, not a failed admin request. """ shape = probe_shape(base_url, api_key, upstream, model) result = ProbeResult( @@ -300,44 +315,22 @@ async def probe_upstream( result.models_error = f"{type(exc).__name__}: {exc}" result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) - started = time.monotonic() if not model_id: return result - request_body = { - "model": model_id, - "messages": [{"role": "user", "content": PROBE_PROMPT}], - "max_tokens": PROBE_MAX_TOKENS, - "stream": False, - } - if endpoint_tag: - request_body["provider"] = { - "order": [endpoint_tag], - "allow_fallbacks": False, - } - try: - async with asyncio.timeout(timeout): - response = await client.post( - result.chat_url, - json=shape_body(request_body, upstream, model), - headers=headers, - params=shape.chat_params, - ) - result.chat_status = response.status_code - result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) - try: - payload = response.json() - except Exception as exc: # noqa: BLE001 - any decode failure is the signal - result.chat_error = f"{type(exc).__name__}: {exc}" - else: - if isinstance(payload, dict): - result.chat_payload = payload - else: - result.chat_error = ( - f"expected a JSON object, got {type(payload).__name__}" - ) - except Exception as exc: # noqa: BLE001 - transport failure is a row status - result.chat_error = f"{type(exc).__name__}: {exc}" - result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) + await _probe_chat( + client, result, model_id, shape, timeout, upstream, model, "max_tokens" + ) + if wants_max_completion_tokens(result.chat_status, result.chat_payload): + await _probe_chat( + client, + result, + model_id, + shape, + timeout, + upstream, + model, + "max_completion_tokens", + ) finally: if owns_client: await client.aclose() @@ -345,6 +338,59 @@ async def probe_upstream( return result +async def _probe_chat( + client: httpx.AsyncClient, + result: ProbeResult, + model_id: str, + shape: ProbeShape, + timeout: float, + upstream: "BaseUpstreamProvider | None", + model: "Model | None", + token_field: str, +) -> None: + """Send the one-token completion and record its outcome on ``result``.""" + request_body: dict[str, Any] = { + "model": model_id, + "messages": [{"role": "user", "content": PROBE_PROMPT}], + token_field: PROBE_MAX_TOKENS, + "stream": False, + } + if result.endpoint_tag: + request_body["provider"] = { + "order": [result.endpoint_tag], + "allow_fallbacks": False, + } + result.token_limit_field = token_field + result.chat_status = None + result.chat_payload = None + result.chat_error = None + started = time.monotonic() + try: + async with asyncio.timeout(timeout): + response = await client.post( + result.chat_url, + json=shape_body(request_body, upstream, model), + headers=shape.headers, + params=shape.chat_params, + ) + result.chat_status = response.status_code + result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) + try: + payload = response.json() + except Exception as exc: # noqa: BLE001 - any decode failure is the signal + result.chat_error = f"{type(exc).__name__}: {exc}" + else: + if isinstance(payload, dict): + result.chat_payload = payload + else: + result.chat_error = ( + f"expected a JSON object, got {type(payload).__name__}" + ) + except Exception as exc: # noqa: BLE001 - transport failure is a row status + result.chat_error = f"{type(exc).__name__}: {exc}" + result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) + + # Row builders are pure: the network lives only in ``probe_upstream`` and # ``run_live_checks``, so every verdict is testable without a socket. @@ -817,13 +863,15 @@ async def run_live_checks( check_cache: bool = True, endpoint_tag: str | None = None, upstream: "BaseUpstreamProvider | None" = None, + advertised_model: "Model | None" = None, ) -> list[dict[str, Any]]: """Probe one upstream and build the live/derived rows. ``check_cache`` adds the prompt-cache and margin rows, which cost two or three more completions against a long prompt. ``upstream`` shapes the probes like the proxy's own requests; without it they assume a plain - OpenAI-compatible base URL. + OpenAI-compatible base URL. ``advertised_model`` carries a pinned path's + own endpoint rates for the margin row to compare against ``model``'s. """ probe = await probe_upstream( base_url, @@ -927,6 +975,8 @@ async def run_live_checks( pricing_known=pricing_known, endpoint_tag=endpoint_tag, upstream=upstream, + advertised_model=advertised_model, + token_limit_field=probe.token_limit_field, ) ) return rows diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index 9c267b69..a7459683 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -99,7 +99,11 @@ class CacheProbeResult: def _request_body( - model_id: str, prefix: str, fmt: str, endpoint_tag: str | None + model_id: str, + prefix: str, + fmt: str, + endpoint_tag: str | None, + token_field: str = "max_tokens", ) -> dict[str, Any]: system: Any if fmt == "cache_control": @@ -118,7 +122,7 @@ def _request_body( {"role": "system", "content": system}, {"role": "user", "content": CACHE_PROBE_QUESTION}, ], - "max_tokens": PROBE_MAX_TOKENS, + token_field: PROBE_MAX_TOKENS, "stream": False, } if endpoint_tag: @@ -184,6 +188,7 @@ async def probe_cache( timeout: float = PROBE_TIMEOUT_SECONDS, upstream: "BaseUpstreamProvider | None" = None, model: "Model | None" = None, + token_limit_field: str = "max_tokens", ) -> CacheProbeResult: """Send the same long prompt twice. @@ -202,7 +207,7 @@ async def probe_cache( async def post( fmt: str, ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: - body = _request_body(model_id, prefix, fmt, endpoint_tag) + body = _request_body(model_id, prefix, fmt, endpoint_tag, token_limit_field) return await _post_completion( http, result.chat_url, @@ -441,6 +446,7 @@ def cost_margin_row( provider_fee: float, sats_to_usd: float, pricing_known: bool = True, + advertised_model: Model | None = None, ) -> dict[str, Any]: """Configured token pricing must cover what the upstream reports charging. @@ -450,11 +456,20 @@ def cost_margin_row( that omit cost, the served ``/v1/models`` list), so a sample where it falls below the fee-adjusted upstream cost means those paths underprice. Upstreams that report no cost give no sample and the row stays a warn. + + ``model`` carries the pricing the proxy reserves and token-bills with. On + a pinned path, ``advertised_model`` carries the endpoint's own rates; a + covered margin whose advertised rates differ from the billed ones is a + warn, since ``/v1/models/paths`` then shows a price the node does not bill. """ + advertised_pricing = ( + advertised_model.sats_pricing if advertised_model is not None else None + ) evidence: dict[str, Any] = { "model_id": model.id, "provider_fee": provider_fee, "sats_usd_price": sats_to_usd, + "pricing_basis": "model pricing the proxy reserves and token-bills with", "samples": [], } if model.sats_pricing is None or not pricing_known: @@ -468,6 +483,7 @@ def cost_margin_row( samples: list[dict[str, Any]] = [] short: list[str] = [] + mismatched: list[str] = [] for payload in payloads: if not isinstance(payload, dict): continue @@ -480,6 +496,11 @@ def cost_margin_row( upstream_total = _expected_usd_msats( reported_usd, provider_fee, sats_to_usd ) + advertised_total = ( + _expected_token_msats(advertised_pricing, usage)[0] + if advertised_pricing is not None + else None + ) except (ValueError, OverflowError) as exc: evidence["error"] = f"{type(exc).__name__}: {exc}" return certification_row( @@ -489,14 +510,17 @@ def cost_margin_row( f"The margin could not be derived: {exc}.", evidence, ) - samples.append( - { - "usage": usage.dict(), - "reported_usd": reported_usd, - "upstream_msats_with_fee": upstream_total, - "configured_msats": configured_total, - } - ) + sample: dict[str, Any] = { + "usage": usage.dict(), + "reported_usd": reported_usd, + "upstream_msats_with_fee": upstream_total, + "configured_msats": configured_total, + } + if advertised_total is not None: + sample["advertised_msats"] = advertised_total + if abs(advertised_total - configured_total) > COST_TOLERANCE_MSATS: + mismatched.append(f"{advertised_total} vs {configured_total}") + samples.append(sample) if configured_total + COST_TOLERANCE_MSATS < upstream_total: short.append(f"{configured_total} < {upstream_total}") evidence["samples"] = samples @@ -521,6 +545,18 @@ def cost_margin_row( "requests lose money.", evidence, ) + if mismatched: + return certification_row( + ROW_MARGIN, + STATUS_WARN, + TITLE_MARGIN, + f"Configured pricing covers the upstream's reported cost on " + f"{len(samples)} sampled completion(s), but this path advertises " + f"different endpoint rates (advertised vs billed msats: " + f"{'; '.join(mismatched)}); the proxy reserves and token-bills " + "pinned requests with the model's own pricing.", + evidence, + ) return certification_row( ROW_MARGIN, STATUS_OK, @@ -561,6 +597,8 @@ async def run_cache_checks( pricing_known: bool = True, endpoint_tag: str | None = None, upstream: "BaseUpstreamProvider | None" = None, + advertised_model: Model | None = None, + token_limit_field: str = "max_tokens", ) -> list[dict[str, Any]]: """Run the cache probe and build the three cache/margin rows.""" probe = await probe_cache( @@ -572,6 +610,7 @@ async def run_cache_checks( timeout=timeout, upstream=upstream, model=model, + token_limit_field=token_limit_field, ) cost_data = await _price_payload(probe.second_payload, model, provider_fee) return [ @@ -595,6 +634,7 @@ async def run_cache_checks( provider_fee=provider_fee, sats_to_usd=sats_to_usd, pricing_known=pricing_known, + advertised_model=advertised_model, ), ), ] diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index db110e53..bcda8d39 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -898,8 +898,8 @@ def apply_model_path_pricing( Direct paths already use the provider model cache and therefore carry the same pricing as ``model``. OpenRouter endpoint rows instead contain raw, - endpoint-specific USD rates; certification must use those rates when its - requests are pinned to that endpoint. + endpoint-specific USD rates; certification compares them against the + model's own pricing, which the proxy reserves and token-bills with. """ if row.endpoint_tag is None: return model diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index a57ce706..9c8b2125 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -328,7 +328,7 @@ async def test_certify_model_path_pins_every_completion( @pytest.mark.integration @pytest.mark.asyncio @respx.mock -async def test_certify_uses_selected_path_pricing_for_margin( +async def test_certify_margin_bills_model_pricing_and_reports_path_pricing( integration_client: AsyncClient, integration_session: AsyncSession ) -> None: base_url = "https://openrouter.ai/api/v1" @@ -409,12 +409,20 @@ async def test_certify_uses_selected_path_pricing_for_margin( assert resp.status_code == 200, resp.text margin = _find_row(resp.json()["rows"], "cost.margin") + # The proxy reserves and token-bills a pinned request with the model's own + # pricing (``configured_msats``); the path's endpoint rates are reported + # alongside (``advertised_msats``) and differ, so the covered margin warns. assert [ - (sample["upstream_msats_with_fee"], sample["configured_msats"]) + ( + sample["upstream_msats_with_fee"], + sample["configured_msats"], + sample["advertised_msats"], + ) for sample in margin["evidence"]["samples"] - ] == [(3, 3), (269, 289), (26, 15)] - assert "289 < 269" not in margin["detail"] - assert "15 < 26" in margin["detail"] + ] == [(3, 3, 3), (269, 356, 289), (26, 43, 15)] + assert margin["status"] == "warn" + assert "advertises different endpoint rates" in margin["detail"] + assert "289 vs 356" in margin["detail"] @pytest.mark.integration diff --git a/tests/integration/test_certify_provider_shapes.py b/tests/integration/test_certify_provider_shapes.py index ee7ab153..c614bc26 100644 --- a/tests/integration/test_certify_provider_shapes.py +++ b/tests/integration/test_certify_provider_shapes.py @@ -88,6 +88,16 @@ SHAPES = [ ] +@pytest.fixture(autouse=True) +def _isolate_proxy_state(monkeypatch: pytest.MonkeyPatch) -> None: + """``reinitialize_upstreams`` rebinds module globals; restore them after + each test so the provider types seeded here never leak into others.""" + from routstr import proxy + + for name in ("_upstreams", "_provider_map", "_unique_models"): + monkeypatch.setattr(proxy, name, getattr(proxy, name).copy()) + + async def _seed(session: AsyncSession, shape: Shape) -> int: provider = UpstreamProviderRow( provider_type=shape.provider_type, @@ -183,3 +193,39 @@ def test_shape_body_keeps_a_single_cache_control_marker() -> None: assert json.dumps(shaped).count('"cache_control"') == 1 assert shaped["model"] == "claude-sonnet-4-5-20250929" + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("shape", SHAPES, ids=[s.provider_type for s in SHAPES]) +async def test_model_test_matches_proxy_request( + shape: Shape, + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """``POST /api/models/test`` reaches the upstream the way the proxy does.""" + with respx.mock(assert_all_called=False) as mock: + await _seed(integration_session, shape) + chat_route = mock.post(shape.chat_url).mock( + return_value=Response(200, json=_mock_chat_response(model=shape.model_id)) + ) + + resp = await integration_client.post( + "/api/models/test", + headers=_admin_headers(), + json={ + "model_id": shape.model_id, + "endpoint_type": "chat-completions", + "request_data": {"messages": [{"role": "user", "content": "hi"}]}, + }, + ) + + assert resp.status_code == 200, resp.text + assert resp.json()["success"] is True, resp.json() + assert chat_route.call_count == 1 + request = chat_route.calls[0].request + header, value = shape.auth_header + assert request.headers.get(header) == value + for key, expected in shape.params.items(): + assert request.url.params.get(key) == expected + assert json.loads(request.content)["model"] == shape.upstream_model diff --git a/tests/integration/test_model_test_endpoint_security.py b/tests/integration/test_model_test_endpoint_security.py index dca7e9fa..b26a84e4 100644 --- a/tests/integration/test_model_test_endpoint_security.py +++ b/tests/integration/test_model_test_endpoint_security.py @@ -196,7 +196,11 @@ async def test_model_test_endpoint_admin_uses_allowed_upstream_path( return None async def post( - self, url: str, json: dict[str, Any], headers: dict[str, str] + self, + url: str, + json: dict[str, Any], + headers: dict[str, str], + params: dict[str, str] | None = None, ) -> MockResponse: assert url == "https://api.example.com/v1/chat/completions" assert json["model"] == "upstream-model-a" @@ -204,7 +208,11 @@ async def test_model_test_endpoint_admin_uses_allowed_upstream_path( return MockResponse() try: - with patch("httpx.AsyncClient", return_value=MockAsyncClient()): + # No live upstream instance: the plain OpenAI-compatible fallback. + with ( + patch("httpx.AsyncClient", return_value=MockAsyncClient()), + patch("routstr.proxy.get_upstreams", return_value=[]), + ): response = await integration_client.post( "/api/models/test", json={ diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py index 2d87a8cf..a1681aeb 100644 --- a/tests/unit/test_certification_cache.py +++ b/tests/unit/test_certification_cache.py @@ -333,6 +333,47 @@ class TestCostMarginRow: assert row["status"] == STATUS_OK + def test_pinned_path_fails_when_billed_pricing_misses_cost(self) -> None: + """The path's endpoint rates cover the cost but the model pricing the + proxy actually bills with does not: the margin must fail.""" + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 4e-6}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + advertised_model=_model(prompt=1e-6, completion=2e-6), + ) + assert row["status"] == STATUS_FAIL + sample = row["evidence"]["samples"][0] + assert sample["advertised_msats"] >= sample["upstream_msats_with_fee"] + assert sample["configured_msats"] < sample["upstream_msats_with_fee"] + + def test_pinned_path_warns_when_advertised_rates_differ(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + advertised_model=_model(prompt=1e-6, completion=2e-6), + ) + assert row["status"] == STATUS_WARN + assert "advertises different endpoint rates" in row["detail"] + + def test_pinned_path_ok_when_advertised_rates_match(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + advertised_model=_model(), + ) + assert row["status"] == STATUS_OK + sample = row["evidence"]["samples"][0] + assert sample["advertised_msats"] == sample["configured_msats"] + def test_warn_when_pricing_unknown(self) -> None: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) row = cost_margin_row( diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py index 7eb05552..b284748c 100644 --- a/tests/unit/test_certification_hardening.py +++ b/tests/unit/test_certification_hardening.py @@ -17,7 +17,6 @@ from typing import Any import pytest -from routstr.payment.usage import parse_token_count from routstr.upstream.certification import ( STATUS_FAIL, STATUS_OK, @@ -50,39 +49,14 @@ def _probe(**kwargs: Any) -> ProbeResult: ) -# Regression: a non-finite token count crashed the billing path. +# Regression: a non-finite token count must not crash the usage row. # ``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``. +# upstream can put them on the wire. Whether the parser rejects them (``warn``, +# nothing to bill on) or raises (``fail``, unreadable usage), the row reports +# it instead of raising. 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( @@ -95,8 +69,7 @@ class TestNonFiniteTokenCounts: }, ) ) - # Both counts collapse to 0, which is the "nothing to bill on" case. - assert row["status"] == STATUS_WARN + assert row["status"] in (STATUS_WARN, STATUS_FAIL) def test_usage_row_survives_infinite_tokens_in_a_string(self) -> None: row = usage_capture_row( @@ -105,7 +78,7 @@ class TestNonFiniteTokenCounts: chat_payload={"usage": {"prompt_tokens": "Infinity"}}, ) ) - assert row["status"] == STATUS_WARN + assert row["status"] in (STATUS_WARN, STATUS_FAIL) # Regression: ``certification_row`` stored non-dict evidence verbatim, so the diff --git a/tests/unit/test_certification_token_limit.py b/tests/unit/test_certification_token_limit.py new file mode 100644 index 00000000..45c8b4da --- /dev/null +++ b/tests/unit/test_certification_token_limit.py @@ -0,0 +1,123 @@ +"""Probes fall back to ``max_completion_tokens`` when ``max_tokens`` is refused. + +OpenAI's o-series and gpt-5 reject ``max_tokens`` on chat completions, so a +probe that only ever sends it fails ``usage.capture`` on a healthy upstream. +""" + +from __future__ import annotations + +import json +from typing import Any + +import httpx +import pytest + +from routstr.upstream.certification import ( + STATUS_OK, + certify_upstream_url, + wants_max_completion_tokens, +) + +OPENAI_REJECTION = { + "error": { + "message": ( + "Unsupported parameter: 'max_tokens' is not supported with this " + "model. Use 'max_completion_tokens' instead." + ), + "type": "invalid_request_error", + "param": "max_tokens", + "code": "unsupported_parameter", + } +} + + +@pytest.fixture(autouse=True) +def _restore_price_globals(monkeypatch: pytest.MonkeyPatch) -> None: + """``certify_upstream_url`` publishes ``sats_usd_price`` to the price + module's globals; restore them so no later test sees this quote.""" + from routstr.payment import price + + monkeypatch.setattr(price, "SATS_USD_PRICE", price.SATS_USD_PRICE) + monkeypatch.setattr(price, "BTC_USD_PRICE", price.BTC_USD_PRICE) + + +def _row(result: dict[str, Any], row_id: str) -> dict[str, Any]: + return next(row for row in result["rows"] if row["id"] == row_id) + + +@pytest.mark.asyncio +async def test_probe_retries_with_max_completion_tokens() -> None: + bodies: list[dict[str, Any]] = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "gpt-5"}]}) + body = json.loads(request.content) + bodies.append(body) + if "max_tokens" in body: + return httpx.Response(400, json=OPENAI_REJECTION) + return httpx.Response( + 200, + json={ + "model": "gpt-5", + "usage": {"prompt_tokens": 8, "completion_tokens": 1}, + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + result = await certify_upstream_url( + "https://mock.example/v1", + model_id="gpt-5", + prompt_price=1e-6, + completion_price=2e-6, + sats_usd_price=0.001, + client=client, + ) + + assert _row(result, "usage.capture")["status"] == STATUS_OK + assert _row(result, "cost.prompt_completion")["status"] == STATUS_OK + # Rejected probe, retried probe, then both cache-probe calls reuse the + # accepted field instead of being rejected again. + assert ["max_tokens" in body for body in bodies] == [True, False, False, False] + assert all(body.get("max_completion_tokens") == 1 for body in bodies[1:]) + + +@pytest.mark.asyncio +async def test_probe_does_not_retry_unrelated_400() -> None: + calls: list[dict[str, Any]] = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "m"}]}) + calls.append(json.loads(request.content)) + return httpx.Response(400, json={"error": {"message": "model not found"}}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + result = await certify_upstream_url( + "https://mock.example/v1", + model_id="m", + prompt_price=1e-6, + completion_price=2e-6, + sats_usd_price=0.001, + client=client, + ) + + assert len(calls) == 1 + assert _row(result, "usage.capture")["status"] != STATUS_OK + + +@pytest.mark.parametrize( + ("status", "payload", "expected"), + [ + (400, OPENAI_REJECTION, True), + (400, {"error": {"message": "bad model"}}, False), + (422, OPENAI_REJECTION, False), + (200, OPENAI_REJECTION, False), + (400, None, False), + (None, None, False), + ], +) +def test_wants_max_completion_tokens( + status: int | None, payload: Any, expected: bool +) -> None: + assert wants_max_completion_tokens(status, payload) is expected From 13976f538b5e5ab34b8f16075ad374a2633d392b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 2 Oct 2026 20:13:11 +0200 Subject: [PATCH 16/18] test: expect warn for non-finite usage now that parsing rejects it --- tests/unit/test_certification_hardening.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py index b284748c..8beb2fce 100644 --- a/tests/unit/test_certification_hardening.py +++ b/tests/unit/test_certification_hardening.py @@ -51,9 +51,8 @@ def _probe(**kwargs: Any) -> ProbeResult: # Regression: a non-finite token count must not crash the usage row. # ``json.loads`` accepts the bare ``Infinity``/``NaN`` literals, so an -# upstream can put them on the wire. Whether the parser rejects them (``warn``, -# nothing to bill on) or raises (``fail``, unreadable usage), the row reports -# it instead of raising. +# upstream can put them on the wire. ``parse_token_count`` treats them as 0, +# which is the "nothing to bill on" ``warn``. class TestNonFiniteTokenCounts: @@ -69,7 +68,7 @@ class TestNonFiniteTokenCounts: }, ) ) - assert row["status"] in (STATUS_WARN, STATUS_FAIL) + assert row["status"] == STATUS_WARN def test_usage_row_survives_infinite_tokens_in_a_string(self) -> None: row = usage_capture_row( @@ -78,7 +77,7 @@ class TestNonFiniteTokenCounts: chat_payload={"usage": {"prompt_tokens": "Infinity"}}, ) ) - assert row["status"] in (STATUS_WARN, STATUS_FAIL) + assert row["status"] == STATUS_WARN # Regression: ``certification_row`` stored non-dict evidence verbatim, so the From 4d65cd6f4c59f05cff131f5917ba7f323722572f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 2 Oct 2026 23:43:00 +0200 Subject: [PATCH 17/18] fix: certify and model-test send model.id like the proxy, not the client alias --- routstr/core/admin.py | 5 +- routstr/payment/models.py | 3 +- routstr/upstream/certification.py | 13 ++- routstr/upstream/certification_cache.py | 2 +- tests/integration/test_certify_alias_paths.py | 4 +- .../test_certify_matches_proxy_model.py | 90 +++++++++++++++++++ 6 files changed, 104 insertions(+), 13 deletions(-) create mode 100644 tests/integration/test_certify_matches_proxy_model.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 957684af..8de51de7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1733,8 +1733,9 @@ async def certify_upstream_provider( provider.provider_fee, sats_to_usd, ) - # The timeout applies per upstream call. The run makes up to five - # calls, so the request can stay open for up to five times it. + # The timeout applies per upstream call. The run makes up to six + # calls (models, two short probes after a max_completion_tokens retry, + # three cache probes), so the request can stay open for six times it. live_rows = await run_live_checks( provider.base_url, provider.api_key, diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 731c6de5..0120e26e 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -689,7 +689,8 @@ def _model_test_target( upstream.build_request_url(path, model_obj), upstream.prepare_headers({"content-type": "application/json"}), dict(upstream.prepare_params(path, None)), - upstream.transform_model_name(model_id), + # The proxy forwards ``model.id``, not the row's client alias. + upstream.transform_model_name(model_obj.id), ) diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index cb4d83e4..fc345f33 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -247,17 +247,14 @@ def shape_body( ) -> Any: """The JSON body the proxy would forward, model-name transforms included. - ``prepare_request_body`` rewrites ``model`` from ``model.id``; the probe - keeps the id it chose (``forwarded_model_id`` first) and only applies the - provider's own name transform to it. + ``prepare_request_body`` sets ``model`` from ``model.id``, exactly as + ``forward_request`` does, so an alias row's ``forwarded_model_id`` never + reaches the upstream here either. """ if upstream is None or model is None: return body shaped = upstream.prepare_request_body(json.dumps(body).encode(), model) - data = json.loads(shaped) if shaped else dict(body) - if isinstance(data, dict) and isinstance(body.get("model"), str): - data["model"] = upstream.transform_model_name(body["model"]) - return data + return json.loads(shaped) if shaped else body async def probe_upstream( @@ -876,7 +873,7 @@ async def run_live_checks( probe = await probe_upstream( base_url, api_key, - model.forwarded_model_id or model.id, + model.id, endpoint_tag=endpoint_tag, client=client, timeout=timeout, diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index a7459683..fef1d8a4 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -604,7 +604,7 @@ async def run_cache_checks( probe = await probe_cache( base_url, api_key, - model.forwarded_model_id or model.id, + model.id, endpoint_tag=endpoint_tag, client=client, timeout=timeout, diff --git a/tests/integration/test_certify_alias_paths.py b/tests/integration/test_certify_alias_paths.py index 6786a8a6..98b35bfd 100644 --- a/tests/integration/test_certify_alias_paths.py +++ b/tests/integration/test_certify_alias_paths.py @@ -106,7 +106,9 @@ async def test_certify_forwarded_alias_listed_path_succeeds( assert response.status_code == 200, response.text assert chat.call_count == 1 body: dict[str, Any] = json.loads(chat.calls[0].request.content) - assert body["model"] == forwarded + # The path is keyed by the exposed id, but the upstream gets what the + # proxy sends for this row: transform_model_name(model.id). + assert body["model"] == "local-alias" assert body["provider"] == {"order": ["endpoint"], "allow_fallbacks": False} mismatch = await integration_client.post( f"/admin/api/upstream-providers/{provider.id}/certify", diff --git a/tests/integration/test_certify_matches_proxy_model.py b/tests/integration/test_certify_matches_proxy_model.py new file mode 100644 index 00000000..d36357bc --- /dev/null +++ b/tests/integration/test_certify_matches_proxy_model.py @@ -0,0 +1,90 @@ +"""Certification must send the upstream the model id the proxy sends. + +An admin alias row has ``id`` and ``forwarded_model_id`` that differ. The +proxy forwards ``transform_model_name(model.id)`` (``prepare_request_body``), +so a probe that sends the forwarded id certifies a request no client can make. +""" + +import json +import time +from typing import Any +from unittest.mock import patch + +import pytest +import respx +from httpx import AsyncClient, Response +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey +from routstr.proxy import reinitialize_upstreams + +from .test_certify_endpoint import _admin_headers, _make_provider, _model_row + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_and_proxy_send_the_same_model( + integration_session: AsyncSession, + integration_client: AsyncClient, +) -> None: + base_url = "https://certify-upstream.example/v1" + respx.get(f"{base_url}/models").mock(return_value=Response(200, json={"data": []})) + chat = respx.post(f"{base_url}/chat/completions").mock( + return_value=Response( + 200, + json={ + "id": "x", + "object": "chat.completion", + "model": "m", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 1}, + }, + ) + ) + provider = await _make_provider(integration_session) + model = _model_row(provider.id, model_id="row-id") # type: ignore[arg-type] + model.forwarded_model_id = "client-alias" + integration_session.add(model) + integration_session.add( + ApiKey( + hashed_key="certify-contract", balance=10**9, created_at=int(time.time()) + ) + ) + await integration_session.commit() + + with ( + patch("routstr.payment.models.sats_usd_price", return_value=0.0005), + patch("routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005), + patch("routstr.payment.price.SATS_USD_PRICE", 0.0005), + ): + await reinitialize_upstreams() + proxied = await integration_client.post( + "/v1/chat/completions", + headers={"Authorization": "Bearer sk-certify-contract"}, + json={ + "model": "client-alias", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 1, + }, + ) + assert proxied.status_code == 200, proxied.text + certified = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={"model_id": "row-id", "check_cache": False}, + ) + assert certified.status_code == 200, certified.text + + bodies: list[dict[str, Any]] = [ + json.loads(call.request.content) for call in chat.calls + ] + assert len(bodies) == 2 + proxy_model, certify_model = bodies[0]["model"], bodies[1]["model"] + assert certify_model == proxy_model From fb7de40a8728288cd4f5fa20f6c5185acdf83fd8 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 3 Oct 2026 13:28:36 +0200 Subject: [PATCH 18/18] fix: serve certification page on direct load, read certify key from env, fix sequential note --- routstr/core/main.py | 27 ++++++----- routstr/upstream/certification.py | 10 +++- tests/unit/test_certification_cli_key.py | 46 +++++++++++++++++++ tests/unit/test_ui_pages_registered.py | 19 ++++++++ .../provider-certification-setup.tsx | 3 +- 5 files changed, 91 insertions(+), 14 deletions(-) create mode 100644 tests/unit/test_certification_cli_key.py create mode 100644 tests/unit/test_ui_pages_registered.py diff --git a/routstr/core/main.py b/routstr/core/main.py index 368202cd..cd4cab90 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -340,6 +340,21 @@ async def providers() -> RedirectResponse: UI_DIST_PATH = Path(__file__).parent.parent.parent / "ui_out" +# Every `ui/app/**/page.tsx` route needs an entry, or a direct load 404s. +UI_PAGES = ( + "dashboard", + "login", + "model", + "providers", + "providers/certification", + "settings", + "transactions", + "balances", + "logs", + "usage", + "unauthorized", +) + if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): logger.info(f"Serving static UI from {UI_DIST_PATH}") @@ -362,18 +377,6 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): # with a slash (e.g. `/login/`). The proxy router catches `/{path:path}` # before FastAPI's `redirect_slashes` logic can normalize the URL, so we # must register both the with-slash and without-slash variants here. - UI_PAGES = ( - "dashboard", - "login", - "model", - "providers", - "settings", - "transactions", - "balances", - "logs", - "usage", - "unauthorized", - ) def _register_ui_page(name: str) -> None: page_dir = UI_DIST_PATH / name diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index fc345f33..180ab239 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -17,6 +17,7 @@ import argparse import asyncio import json import math +import os import sys import time from collections.abc import Callable @@ -1245,7 +1246,14 @@ def main(argv: list[str] | None = None) -> int: required=True, help="Upstream base URL (repeatable), e.g. https://api.example.com/v1", ) - parser.add_argument("--key", default="", help="Bearer API key for the upstream") + parser.add_argument( + "--key", + default=os.environ.get("ROUTSTR_CERTIFY_KEY", ""), + help=( + "Bearer API key for the upstream (defaults to $ROUTSTR_CERTIFY_KEY; " + "prefer the env var so the key stays out of shell history and ps)" + ), + ) parser.add_argument( "--model", default=None, diff --git a/tests/unit/test_certification_cli_key.py b/tests/unit/test_certification_cli_key.py new file mode 100644 index 00000000..c9355445 --- /dev/null +++ b/tests/unit/test_certification_cli_key.py @@ -0,0 +1,46 @@ +"""The certification CLI reads the upstream key from the environment.""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from routstr.upstream import certification + + +@pytest.fixture +def captured_keys(monkeypatch: pytest.MonkeyPatch) -> list[str]: + keys: list[str] = [] + + async def fake_certify(url: str, *, api_key: str, **_: Any) -> dict[str, Any]: + keys.append(api_key) + return {"url": url, "rows": []} + + monkeypatch.setattr(certification, "certify_upstream_url", fake_certify) + monkeypatch.setattr(certification, "render_checklist", lambda _result: "") + return keys + + +def test_key_defaults_to_env_var( + monkeypatch: pytest.MonkeyPatch, captured_keys: list[str] +) -> None: + monkeypatch.setenv("ROUTSTR_CERTIFY_KEY", "sk-from-env") + assert certification.main(["--url", "http://localhost:1/v1"]) == 0 + assert captured_keys == ["sk-from-env"] + + +def test_key_flag_overrides_env_var( + monkeypatch: pytest.MonkeyPatch, captured_keys: list[str] +) -> None: + monkeypatch.setenv("ROUTSTR_CERTIFY_KEY", "sk-from-env") + certification.main(["--url", "http://localhost:1/v1", "--key", "sk-flag"]) + assert captured_keys == ["sk-flag"] + + +def test_key_is_empty_without_flag_or_env( + monkeypatch: pytest.MonkeyPatch, captured_keys: list[str] +) -> None: + monkeypatch.delenv("ROUTSTR_CERTIFY_KEY", raising=False) + certification.main(["--url", "http://localhost:1/v1"]) + assert captured_keys == [""] diff --git a/tests/unit/test_ui_pages_registered.py b/tests/unit/test_ui_pages_registered.py new file mode 100644 index 00000000..3ccdd3a8 --- /dev/null +++ b/tests/unit/test_ui_pages_registered.py @@ -0,0 +1,19 @@ +"""Every static UI page must be served on a direct load, not the proxy 404.""" + +from __future__ import annotations + +from pathlib import Path + +from routstr.core import main as core_main + +UI_APP_DIR = Path(__file__).resolve().parents[2] / "ui" / "app" + + +def test_every_ui_app_page_is_in_ui_pages() -> None: + routes = { + page.parent.relative_to(UI_APP_DIR).as_posix() + for page in UI_APP_DIR.rglob("page.tsx") + if page.parent != UI_APP_DIR + } + assert routes, "no ui/app pages found" + assert routes - set(core_main.UI_PAGES) == set() diff --git a/ui/components/provider-certification-setup.tsx b/ui/components/provider-certification-setup.tsx index 6c2f0a98..d8ada6bb 100644 --- a/ui/components/provider-certification-setup.tsx +++ b/ui/components/provider-certification-setup.tsx @@ -272,7 +272,8 @@ export function ProviderCertificationSetupPanel({ )} {mode === 'all' && (

    - All {paths.length} paths will run in parallel. + All {paths.length} paths will run one after another, each + with its own probe calls.

    )}