diff --git a/pyproject.toml b/pyproject.toml index bd6f4550..9868e774 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,6 +38,7 @@ dev = [ "psutil>=5.9.0", "aiohttp>=3.9.0", "pytest-benchmark>=4.0.0", + "respx>=0.21", "routstr", ] diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 46406381..2041f409 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 ( @@ -24,6 +28,7 @@ from .db import ( CashuTransaction, CliToken, LightningInvoice, + ModelPathRow, ModelRow, UpstreamProviderRow, create_session, @@ -1234,6 +1239,27 @@ 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 exposed_model_id + + certification_paths: dict[str, list[dict[str, object]]] = {} + for model in [*db_models, *filtered_remote_models]: + paths = paths_by_public_id.get(exposed_model_id(model).lower(), []) + certification_paths[model.id] = paths + return { "provider": { "id": provider.id, @@ -1247,9 +1273,495 @@ 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, } +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 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 + 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: + # 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 + + 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 CertifyRequest(BaseModel): + model_id: str | None = None + model_path: str | None = None + check_cache: bool = True + + +@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 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 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 + :mod:`routstr.upstream.certification`, and a ``checklist`` of the + operator-facing goals. + """ + from ..upstream.certification import ( + PROBE_TIMEOUT_SECONDS, + 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()) + + 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 + + selector = decode_model_path(payload.model_path) + if selector is None: + raise HTTPException(status_code=400, detail="Malformed model path") + 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 + path_model_id = selector.model_id + + 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: + # 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 + break + if model_id is None: + model_id = enabled_rows[0].id + + 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: + 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 + + # 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 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 + + 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 path_model_id is None + or selected_id.lower() != path_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, + certification_row, + ) + from ..upstream.certification_cache import skipped_cache_rows + + 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.", + {}, + ), + *skipped_cache_rows("Skipped — no model to probe."), + ] + else: + 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", + ) + # 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 + + advertised_model = apply_model_path_pricing( + model_obj, + selected_path, + provider.provider_fee, + sats_to_usd, + ) + # 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, + model_obj, + provider_fee=provider.provider_fee, + sats_to_usd=sats_to_usd, + timeout=PROBE_TIMEOUT_SECONDS, + check_cache=payload.check_cache, + endpoint_tag=endpoint_tag, + upstream=upstream_obj, + advertised_model=advertised_model, + ) + + 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/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/payment/models.py b/routstr/payment/models.py index 3f925e19..cf1bc07a 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -657,6 +657,45 @@ 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)), + # The proxy forwards ``model.id``, not the row's client alias. + upstream.transform_model_name(model_obj.id), + ) + + @models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) async def test_model( payload: ModelTestRequest, @@ -690,8 +729,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")) @@ -700,9 +741,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={ @@ -714,14 +752,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/upstream/certification.py b/routstr/upstream/certification.py new file mode 100644 index 00000000..180ab239 --- /dev/null +++ b/routstr/upstream/certification.py @@ -0,0 +1,1361 @@ +"""Live certification checks for an upstream provider endpoint. + +Extends the read-only pricing rows, which never touch the network, with the +ones that must: a ``/models`` heartbeat and a one-token completion. + +Probes call the upstream directly with ``httpx``, never through the node's +billing path — no reservation, no Cashu. Upstream spend is one one-token +completion, plus two or three one-token completions on a ~4.4k-token prompt +when the cache checks are enabled (per certified model path). +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 + +import argparse +import asyncio +import json +import math +import os +import sys +import time +from collections.abc import Callable +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 _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__) + +STATUS_OK = "ok" +STATUS_WARN = "warn" +STATUS_FAIL = "fail" + +TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"} + +# Bounded so a dead upstream fails the row rather than wedging the request. +PROBE_TIMEOUT_SECONDS = 15.0 + +# The cheapest request that still exercises the usage/cost path. +PROBE_MAX_TOKENS = 1 +PROBE_PROMPT = "ping" + +# ``calculate_cost`` demands a reservation ceiling; any value at or above the +# real charge behaves identically. +_PROBE_MAX_COST_MSATS = 1_000_000_000 + +# ``_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 + + +def certification_row( + row_id: str, + status: str, + title: str, + detail: str, + evidence: dict[str, Any] | None = None, +) -> dict[str, Any]: + """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, + "title": title, + "detail": detail, + "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, 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 + 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}, + ) + + +# 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", + "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"), + ), + ( + "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",), + ), +) + + +def build_checklist(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + 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 + endpoint_tag: str | None = None + 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 + # ``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 +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`` 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) + return json.loads(shaped) if shaped else body + + +async def probe_upstream( + base_url: str, + api_key: str, + model_id: str, + *, + 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. + + 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( + base_url=base_url, + models_url=shape.models_url, + chat_url=shape.chat_url, + endpoint_tag=endpoint_tag, + ) + headers = shape.headers + + owns_client = client is None + if client is None: + client = httpx.AsyncClient(timeout=timeout) + + try: + started = time.monotonic() + try: + async with asyncio.timeout(timeout): + 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: + 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) + + if not model_id: + return result + await _probe_chat( + client, result, model_id, shape, timeout, upstream, model, "max_tokens" + ) + if wants_max_completion_tokens(result.chat_status, result.chat_payload): + await _probe_chat( + client, + result, + model_id, + shape, + timeout, + upstream, + model, + "max_completion_tokens", + ) + finally: + if owns_client: + await client.aclose() + + 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. + + +def endpoint_validity_row(base_url: str) -> dict[str, Any]: + 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``; + # only ``.hostname`` answers whether there is a host to connect to. + if not parsed.hostname: + problems.append("no host component") + evidence: dict[str, Any] = { + "base_url": base_url, + "scheme": parsed.scheme, + "host": parsed.hostname, + "port": parsed.port, + "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]: + 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]: + 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 or type(payload).__name__}.", + {"url": probe.models_url, "error": probe.models_error}, + ) + + data = 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(payload.keys()), + }, + ) + + ids = [ + item["id"] + for item in data + if isinstance(item, dict) and isinstance(item.get("id"), str) and item["id"] + ] + 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 non-empty 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 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, + "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 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: " + f"{probe.chat_error or type(probe.chat_payload).__name__}.", + evidence, + ) + + raw_usage = probe.chat_payload.get("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( + "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. + + 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 + return _resolve_usd_cost(usage, payload) + + +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. + + 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`` + on a non-finite rate, which would otherwise crash ``math.ceil`` downstream. + """ + 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) + 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.""" + 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) + + +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 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 + + 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, + "provider_fee": provider_fee, + "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( + "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, + ) + + reported_usd = _reported_usd_cost(payload) + 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, + ) + + 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, + 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. ``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, + api_key, + model.id, + endpoint_tag=endpoint_tag, + client=client, + timeout=timeout, + upstream=upstream, + model=model, + ) + rows = [ + 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 + 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( + 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, + ), + ) + ) + + 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 + 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.") + ) + 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, + upstream=upstream, + advertised_model=advertised_model, + token_limit_field=probe.token_limit_field, + ) + ) + return rows + + +# 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: + 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: + """A USD-per-token price from outside the node, or ``None``. + + Shares ``coerce_rate`` so an explicit ``--prompt-price`` is validated + exactly like a litellm-derived one. + """ + 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. + """ + 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: + 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() + 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 + + +def _model_from_usd_pricing( + 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 ( + 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 * 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, + 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, + 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 + + 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: + 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.", + {}, + ), + ] + from .certification_cache import skipped_cache_rows + + rows.extend(skipped_cache_rows("Skipped — no model to probe.")) + return { + "target": target, + "rows": rows, + "checklist": build_checklist(rows), + } + + entry = litellm_cost_entry(model_id) or {} + resolved_prompt = ( + _as_price(prompt_price) + if prompt_price is not None + else _as_price(entry.get("input_cost_per_token")) + ) + resolved_completion = ( + _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 + 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 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, + api_key, + model, + provider_fee=provider_fee, + sats_to_usd=sats_to_usd or 1.0, + client=client, + timeout=timeout, + pricing_known=pricing_known, + check_cache=check_cache, + ) + return {"target": target, "rows": rows, "checklist": build_checklist(rows)} + + +def render_checklist(result: dict[str, Any]) -> str: + 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 _route_logs_to_stderr() -> None: + """Move the app's stdout log handlers to stderr, so log records cannot + interleave with the report.""" + 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=( + "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=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, + 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( + "--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, + 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" + ) + 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." + ), + ) + 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]]: + 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, + sats_usd_price=args.sats_usd_price, + check_cache=not args.no_cache, + ) + ) + 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)) + elif args.json_out is None: + 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/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py new file mode 100644 index 00000000..fef1d8a4 --- /dev/null +++ b/routstr/upstream/certification_cache.py @@ -0,0 +1,648 @@ +"""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 asyncio +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, + _fixed_token_pricing_active, + _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__) + +# 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, + token_field: str = "max_tokens", +) -> 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}, + ], + token_field: 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: Any, + headers: dict[str, str], + params: dict[str, str], + timeout: float, +) -> tuple[int | None, dict[str, Any] | None, str | None, float]: + started = time.monotonic() + try: + async with asyncio.timeout(timeout): + response = await client.post(url, json=body, headers=headers, params=params) + except Exception as exc: # noqa: BLE001 - transport failure is a row status + latency = round((time.monotonic() - started) * 1000, 2) + return None, None, f"{type(exc).__name__}: {exc}", latency + latency = round((time.monotonic() - started) * 1000, 2) + try: + payload = response.json() + except Exception as exc: # noqa: BLE001 - any decode failure is the signal + return response.status_code, None, f"{type(exc).__name__}: {exc}", latency + if not isinstance(payload, dict): + return ( + response.status_code, + None, + f"expected a JSON object, got {type(payload).__name__}", + latency, + ) + return response.status_code, payload, None, latency + + +def _record( + result: CacheProbeResult, + outcome: tuple[int | None, dict[str, Any] | None, str | None, float], +) -> None: + status, payload, error, latency = outcome + result.statuses.append(status) + result.payloads.append(payload) + result.errors.append(error) + result.latencies_ms.append(latency) + + +def _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, + upstream: "BaseUpstreamProvider | None" = None, + model: "Model | None" = None, + token_limit_field: str = "max_tokens", +) -> 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 on HTTP 400/422, and the second call mirrors whichever format + succeeded. Each call's elapsed deadline includes the response body. + """ + shape = probe_shape(base_url, api_key, upstream, model) + result = CacheProbeResult(chat_url=shape.chat_url, endpoint_tag=endpoint_tag) + prefix = cache_probe_prefix() + + owns_client = client is None + http = client if client is not None else httpx.AsyncClient(timeout=timeout) + + async def post( + fmt: str, + ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: + body = _request_body(model_id, prefix, fmt, endpoint_tag, token_limit_field) + return await _post_completion( + http, + result.chat_url, + 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("plain") + _record(result, first) + if not _is_2xx(first[0]): + return result + second = await post(result.request_format) + _record(result, second) + finally: + if owns_client: + await http.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, + ) + + 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( + 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 + input_rate, _, cache_read_rate, _ = _token_rates(pricing) + 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_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, + "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: + 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} " + f"msats) because {reason}; 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, + advertised_model: Model | None = None, +) -> 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. + + ``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: + 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] = [] + mismatched: 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 + ) + 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( + ROW_MARGIN, + STATUS_FAIL, + TITLE_MARGIN, + f"The margin could not be derived: {exc}.", + evidence, + ) + 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 + + 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, + ) + 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, + 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, + 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( + base_url, + api_key, + model.id, + endpoint_tag=endpoint_tag, + client=client, + timeout=timeout, + upstream=upstream, + model=model, + token_limit_field=token_limit_field, + ) + 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, + advertised_model=advertised_model, + ), + ), + ] + + +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 a53ec0f0..bcda8d39 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -38,6 +38,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__) @@ -887,6 +888,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 compares them against the + model's own pricing, which the proxy reserves and token-bills with. + """ + 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_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"]) diff --git a/tests/integration/test_certify_alias_paths.py b/tests/integration/test_certify_alias_paths.py new file mode 100644 index 00000000..98b35bfd --- /dev/null +++ b/tests/integration/test_certify_alias_paths.py @@ -0,0 +1,133 @@ +"""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 .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) + # 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", + 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/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py new file mode 100644 index 00000000..9c8b2125 --- /dev/null +++ b/tests/integration/test_certify_endpoint.py @@ -0,0 +1,1029 @@ +"""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 ModelPathRow, ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams +from routstr.upstream.generic import GenericUpstreamProvider +from routstr.upstream.model_paths import encode_model_path + + +# 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 + ) + assert provider.id is not None + 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, + }, + } + + +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( + 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, + 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( + side_effect=_caching_upstream() + ) + + 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 + + live_row_ids = [ + "endpoint.validity", + "endpoint.reachable", + "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) + assert row["status"] == "ok", f"{row_id}: {row}" + + 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_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_margin_bills_model_pricing_and_reports_path_pricing( + 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", 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") + # 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["advertised_msats"], + ) + for sample in margin["evidence"]["samples"] + ] == [(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 +@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 +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 +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 = 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()) + ) + 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) + 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( + 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() + 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_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(GenericUpstreamProvider): + 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")) + ) + + 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(), + 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 +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 {"☑️", "⚠️", "❌"} + + +@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/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 diff --git a/tests/integration/test_certify_provider_shapes.py b/tests/integration/test_certify_provider_shapes.py new file mode 100644 index 00000000..c614bc26 --- /dev/null +++ b/tests/integration/test_certify_provider_shapes.py @@ -0,0 +1,231 @@ +"""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"), + ), +] + + +@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, + 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" + + +@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.py b/tests/unit/test_certification.py new file mode 100644 index 00000000..8c1d82e4 --- /dev/null +++ b/tests/unit/test_certification.py @@ -0,0 +1,729 @@ +"""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 + +import pytest + +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), + 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) == 6 + 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 + + +# --- 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 new file mode 100644 index 00000000..a1681aeb --- /dev/null +++ b/tests/unit/test_certification_cache.py @@ -0,0 +1,566 @@ +"""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: list[dict[str, Any] | None] = [ + _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: list[dict[str, Any] | None] = [ + _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_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( + 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"] + + @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( + 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_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_certification_hardening.py b/tests/unit/test_certification_hardening.py new file mode 100644 index 00000000..8beb2fce --- /dev/null +++ b/tests/unit/test_certification_hardening.py @@ -0,0 +1,339 @@ +"""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.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"), + ) + + +# 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. ``parse_token_count`` treats them as 0, +# which is the "nothing to bill on" ``warn``. + + +class TestNonFiniteTokenCounts: + 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"), + } + }, + ) + ) + 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 + + +# Regression: ``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} + + +# Regression: ``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 + + +# 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: + @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 + + +# 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: + 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 + + +# Regression: 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) + + +# Regression: ``_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 + + +# Regression: 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 + + +# Regression: 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) + + +# 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]: + 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"]) == 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") + 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 diff --git a/tests/unit/test_certification_review_regressions.py b/tests/unit/test_certification_review_regressions.py new file mode 100644 index 00000000..92632b45 --- /dev/null +++ b/tests/unit/test_certification_review_regressions.py @@ -0,0 +1,310 @@ +"""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: 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", + 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/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 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/app/providers/certification/page.tsx b/ui/app/providers/certification/page.tsx new file mode 100644 index 00000000..93274331 --- /dev/null +++ b/ui/app/providers/certification/page.tsx @@ -0,0 +1,594 @@ +'use client'; + +import { useEffect, useMemo, useRef, 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, + getSelectedCertificationResults, +} 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 activeRuns = useRef(new Set()); + const generation = useRef(0); + + useEffect(() => { + return () => { + generation.current += 1; + }; + }, []); + + 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) => { + 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]: [] })); + setProgressByProvider((current) => ({ ...current, [providerId]: null })); + setRunningProviderIds((current) => + current.includes(providerId) ? current : [...current, providerId] + ); + try { + await runProviderCertification({ + providerId, + modelRuns, + includeCache: setup.checkCache, + shouldContinue, + onProgress: (progress) => { + if (shouldContinue()) { + setProgressByProvider((current) => ({ + ...current, + [providerId]: progress, + })); + } + }, + onResults: (results) => { + if (shouldContinue()) { + setResultsByProvider((current) => ({ + ...current, + [providerId]: results, + })); + } + }, + }); + } finally { + activeRuns.current.delete(providerId); + if (shouldContinue()) { + 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 = getSelectedCertificationResults( + selectedProviderIds, + resultsByProvider + ); + 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={ - - - + + + + } /> + + + + + + + + + + + + ); +} diff --git a/ui/components/provider-certification-results.tsx b/ui/components/provider-certification-results.tsx new file mode 100644 index 00000000..8fa7cd14 --- /dev/null +++ b/ui/components/provider-certification-results.tsx @@ -0,0 +1,305 @@ +'use client'; + +import { useEffect, useState } from 'react'; +import { + AlertTriangle, + CheckCircle2, + ChevronDown, + ChevronRight, + Loader2, + XCircle, +} from 'lucide-react'; + +import { Badge } from '@/components/ui/badge'; +import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; +import type { + CertificationRow, + CertificationStatus, + ProviderCertification, +} from '@/lib/api/services/admin'; +import type { + CertificationProgress, + ModelCertificationResult, +} from '@/lib/provider-certification'; +import { getCertificationResultStatus } from '@/lib/provider-certification'; +import { cn } from '@/lib/utils'; + +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', + }, +}; + +export interface CertificationResultSummary { + ok: number; + warn: number; + fail: number; + error: number; +} + +export function summarizeCertificationResults( + results: ModelCertificationResult[] +): CertificationResultSummary { + const summary = { ok: 0, warn: 0, fail: 0, error: 0 }; + for (const result of results) { + summary[getCertificationResultStatus(result)] += 1; + } + return summary; +} + +function StatusBadge({ status }: { status: CertificationStatus }) { + const style = STATUS_STYLES[status]; + const Icon = style.icon; + return ( + + + {style.label} + + ); +} + +function ResultStatusBadge({ + status, +}: { + status: CertificationStatus | 'error'; +}) { + if (status === 'error') { + return ( + + Error + + ); + } + return ; +} + +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 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..d8ada6bb --- /dev/null +++ b/ui/components/provider-certification-setup.tsx @@ -0,0 +1,316 @@ +'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 one after another, each + with its own probe calls. +

    + )} + + ) : ( +

    + 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} + /> + +
    + {setup.checkCache && ( +

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

    + )} +
    + ); +} 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 new file mode 100644 index 00000000..8bd117aa --- /dev/null +++ b/ui/hooks/use-provider-certification-runner.ts @@ -0,0 +1,141 @@ +'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; + shouldContinue?: () => boolean; + onProgress?: (progress: CertificationProgress | null) => void; + onResults?: (results: ModelCertificationResult[]) => void; +} + +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, + modelTotal: modelRuns.length, + pathCount: run.targets.length, + }); + // 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) { + if (!shouldContinue()) break; + 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), + }); + } + } + if (!shouldContinue()) break; + completed.push(...batch); + onResults?.([...completed]); + } + + if (shouldContinue()) 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); + const active = useRef(false); + + useEffect(() => { + mounted.current = true; + return () => { + mounted.current = false; + generation.current += 1; + }; + }, []); + + const reset = useCallback(() => { + generation.current += 1; + setResults([]); + setProgress(null); + // 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([]); + setProgress(null); + setIsPending(true); + try { + return await runProviderCertification({ + providerId, + modelRuns, + includeCache, + shouldContinue: () => + mounted.current && generation.current === runGeneration, + onProgress: (nextProgress) => { + if (mounted.current && generation.current === runGeneration) { + setProgress(nextProgress); + } + }, + onResults: (nextResults) => { + if (mounted.current && generation.current === runGeneration) { + setResults(nextResults); + } + }, + }); + } finally { + active.current = false; + if (mounted.current) { + setProgress(null); + setIsPending(false); + } + } + }, + [providerId] + ); + + return { results, progress, isPending, run, reset }; +} diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index 15daa38d..52ec852f 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -46,6 +46,42 @@ 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; + check_cache?: boolean; +}; + export const AdminModelPricingSchema = z.object({ prompt: z.number().optional(), completion: z.number().optional(), @@ -84,6 +120,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 +134,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)), }); export type ProviderType = z.infer; @@ -107,6 +150,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 +361,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` diff --git a/ui/lib/provider-certification.ts b/ui/lib/provider-certification.ts new file mode 100644 index 00000000..dff4478c --- /dev/null +++ b/ui/lib/provider-certification.ts @@ -0,0 +1,123 @@ +import { getApiErrorMessage } from '@/lib/api/errors'; +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 getSelectedCertificationResults = ( + providerIds: number[], + resultsByProvider: Record +): ModelCertificationResult[] => + providerIds.flatMap((providerId) => resultsByProvider[providerId] ?? []); + +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 => + getApiErrorMessage(error, 'Certification request failed'); diff --git a/uv.lock b/uv.lock index 1764054b..6abbc6ab 100644 --- a/uv.lock +++ b/uv.lock @@ -2688,6 +2688,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" @@ -2751,6 +2763,7 @@ dev = [ { name = "pytest-asyncio" }, { name = "pytest-benchmark" }, { name = "pytest-cov" }, + { name = "respx" }, { name = "routstr" }, { name = "ruff" }, ] @@ -2788,6 +2801,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" }, ]