Merge pull request #760 from 9qeklajc/feat/upstream-certification-harness

feat(certification): comprehensive upstream certification harness with live probes
This commit is contained in:
9qeklajc
2026-10-03 13:56:23 +02:00
committed by GitHub
31 changed files with 8955 additions and 32 deletions
+1
View File
@@ -38,6 +38,7 @@ dev = [
"psutil>=5.9.0",
"aiohttp>=3.9.0",
"pytest-benchmark>=4.0.0",
"respx>=0.21",
"routstr",
]
+513 -1
View File
@@ -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
+15 -12
View File
@@ -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
+46 -11
View File
@@ -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:
File diff suppressed because it is too large Load Diff
+648
View File
@@ -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, {}),
]
+52
View File
@@ -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:
@@ -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"])
@@ -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
File diff suppressed because it is too large Load Diff
@@ -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
@@ -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
@@ -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={
+729
View File
@@ -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
+566
View File
@@ -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"]
+46
View File
@@ -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 == [""]
+339
View File
@@ -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
@@ -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 []
)
@@ -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
+19
View File
@@ -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()
+594
View File
@@ -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<number[]>([]);
const [activeProviderId, setActiveProviderId] = useState<number | null>(null);
const [setups, setSetups] = useState<
Record<number, ProviderCertificationSetup>
>({});
const [workspaceTab, setWorkspaceTab] = useState<'setup' | 'results'>(
'setup'
);
const [resultsByProvider, setResultsByProvider] = useState<
Record<number, ModelCertificationResult[]>
>({});
const [progressByProvider, setProgressByProvider] = useState<
Record<number, CertificationProgress | null>
>({});
const [runningProviderIds, setRunningProviderIds] = useState<number[]>([]);
const activeRuns = useRef(new Set<number>());
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 (
<AppPageShell contentClassName='mx-auto flex w-full max-w-7xl flex-col'>
<div className='flex min-h-0 flex-1 flex-col gap-4'>
<PageHeader
title='Provider Certification'
description='Configure and certify models across multiple upstream providers.'
actions={
<div className='flex gap-2'>
<Button asChild variant='outline'>
<Link href='/providers'>
<ArrowLeft className='h-4 w-4' />
Providers
</Link>
</Button>
<Popover>
<PopoverTrigger asChild>
<Button variant='outline'>
<Server className='h-4 w-4' />
Select providers
<Badge variant='secondary'>
{selectedProviderIds.length}
</Badge>
<ChevronDown className='h-4 w-4' />
</Button>
</PopoverTrigger>
<PopoverContent align='end' className='w-80 p-0'>
<Command className='h-auto'>
<CommandInput placeholder='Search providers…' />
<CommandList className='max-h-72'>
<CommandEmpty>
{providersQuery.isLoading
? 'Loading providers…'
: 'No providers found'}
</CommandEmpty>
{providers.map((provider) => (
<CommandItem
key={provider.id}
value={`${providerName(provider)} ${provider.base_url}`}
onSelect={() => toggleProvider(provider.id)}
disabled={runningProviderIds.includes(provider.id)}
>
<Checkbox
checked={selectedProviderIds.includes(provider.id)}
tabIndex={-1}
aria-hidden='true'
className='pointer-events-none'
/>
<div className='min-w-0 flex-1'>
<div className='truncate'>
{providerName(provider)}
</div>
<div className='text-muted-foreground truncate text-xs'>
{provider.base_url}
</div>
</div>
<Badge
variant={provider.enabled ? 'secondary' : 'outline'}
className='text-[10px]'
>
{provider.enabled ? 'Enabled' : 'Disabled'}
</Badge>
</CommandItem>
))}
</CommandList>
</Command>
</PopoverContent>
</Popover>
</div>
}
/>
{selectedProviders.length === 0 ? (
<Card className='flex min-h-72 items-center justify-center'>
<CardContent className='space-y-3 text-center'>
<Server className='text-muted-foreground mx-auto h-8 w-8' />
<div>
<div className='font-medium'>Select providers to certify</div>
<p className='text-muted-foreground mt-1 text-sm'>
Choose two or more providers to configure independent model
and path runs.
</p>
</div>
</CardContent>
</Card>
) : (
<div className='grid min-h-0 flex-1 gap-4 md:grid-cols-[240px_minmax(0,1fr)]'>
<aside className='hidden min-h-0 space-y-2 overflow-y-auto rounded-lg border p-2 md:block'>
{selectedProviders.map((provider) => {
const setup = setups[provider.id] ?? emptyCertificationSetup();
const runs = providerRuns(provider.id);
const results = resultsByProvider[provider.id] ?? [];
const summary = summarizeCertificationResults(results);
const running = runningProviderIds.includes(provider.id);
return (
<button
key={provider.id}
type='button'
onClick={() => setActiveProviderId(provider.id)}
className={cn(
'hover:bg-muted w-full space-y-2 rounded-md border p-3 text-left transition-colors',
activeProviderId === provider.id &&
'border-primary bg-muted/60'
)}
>
<div className='flex items-start justify-between gap-2'>
<div className='min-w-0'>
<div className='truncate text-sm font-medium'>
{providerName(provider)}
</div>
<div className='text-muted-foreground truncate text-xs'>
{provider.base_url}
</div>
</div>
{running ? (
<Loader2 className='h-4 w-4 shrink-0 animate-spin' />
) : results.length > 0 ? (
<CheckCircle2 className='h-4 w-4 shrink-0 text-emerald-600' />
) : null}
</div>
<div className='text-muted-foreground flex flex-wrap gap-1 text-[11px]'>
<span>{setup.selectedModelIds.length} models</span>
<span>·</span>
<span>{countCertificationTargets(runs)} routes</span>
</div>
{results.length > 0 && (
<div className='flex flex-wrap gap-1'>
<Badge variant='secondary' className='text-[10px]'>
{summary.ok} ok
</Badge>
{(summary.warn > 0 ||
summary.fail > 0 ||
summary.error > 0) && (
<Badge variant='outline' className='text-[10px]'>
{summary.warn + summary.fail + summary.error} issues
</Badge>
)}
</div>
)}
</button>
);
})}
</aside>
<div className='min-h-0 space-y-3 md:hidden'>
<Select
value={activeProviderId?.toString()}
onValueChange={(value) => setActiveProviderId(Number(value))}
>
<SelectTrigger className='w-full'>
<SelectValue placeholder='Choose active provider' />
</SelectTrigger>
<SelectContent>
{selectedProviders.map((provider) => (
<SelectItem
key={provider.id}
value={provider.id.toString()}
>
{providerName(provider)}
</SelectItem>
))}
</SelectContent>
</Select>
</div>
{activeProvider && activeSetup ? (
<Card className='flex min-h-[36rem] min-w-0 flex-col overflow-hidden p-0 md:min-h-0'>
<div className='shrink-0 border-b px-4 py-3'>
<div className='flex flex-wrap items-start justify-between gap-2'>
<div className='min-w-0'>
<div className='font-medium'>
{providerName(activeProvider)}
</div>
<div className='text-muted-foreground truncate text-xs'>
{activeProvider.base_url}
</div>
</div>
<Badge
variant={activeProvider.enabled ? 'secondary' : 'outline'}
>
{activeProvider.enabled ? 'Enabled' : 'Disabled'}
</Badge>
</div>
</div>
<Tabs
value={workspaceTab}
onValueChange={(value) =>
setWorkspaceTab(value as 'setup' | 'results')
}
className='min-h-0 flex-1 overflow-hidden px-4 pb-4'
>
<TabsList className='grid w-full shrink-0 grid-cols-2'>
<TabsTrigger value='setup'>Setup</TabsTrigger>
<TabsTrigger value='results'>
Results
{activeResults.length > 0 && (
<Badge
variant='secondary'
className='ml-1 px-1.5 py-0 text-xs'
>
{activeResults.length}
</Badge>
)}
</TabsTrigger>
</TabsList>
<TabsContent
value='setup'
className='mt-0 min-h-0 overflow-hidden data-[state=active]:flex data-[state=active]:flex-col'
>
<div className='min-h-0 flex-1 overflow-y-auto pr-1'>
<ProviderCertificationSetupPanel
models={
activeModelsQuery?.data as ProviderModels | undefined
}
isLoading={activeModelsQuery?.isLoading}
error={activeModelsQuery?.error}
setup={activeSetup}
onChange={(next) =>
updateProviderSetup(activeProvider.id, next)
}
disabled={runningProviderIds.includes(
activeProvider.id
)}
idPrefix={`multi-certify-${activeProvider.id}`}
/>
</div>
<div className='mt-3 flex shrink-0 justify-end border-t pt-3'>
<Button
variant='outline'
size='sm'
onClick={() => {
setWorkspaceTab('results');
void runOneProvider(activeProvider.id);
}}
disabled={
runningProviderIds.includes(activeProvider.id) ||
!isProviderReady(activeProvider.id)
}
>
{runningProviderIds.includes(activeProvider.id) ? (
<Loader2 className='h-4 w-4 animate-spin' />
) : (
<RotateCcw className='h-4 w-4' />
)}
Run this provider
</Button>
</div>
</TabsContent>
<TabsContent
value='results'
className='mt-0 min-h-0 overflow-hidden data-[state=active]:flex data-[state=active]:flex-col'
>
<div className='mb-2 flex shrink-0 justify-end'>
<Button
variant='outline'
size='sm'
onClick={() => void runOneProvider(activeProvider.id)}
disabled={
runningProviderIds.includes(activeProvider.id) ||
!isProviderReady(activeProvider.id)
}
>
{runningProviderIds.includes(activeProvider.id) ? (
<Loader2 className='h-4 w-4 animate-spin' />
) : (
<RotateCcw className='h-4 w-4' />
)}
Run this provider again
</Button>
</div>
<ProviderCertificationResults
results={activeResults}
progress={progressByProvider[activeProvider.id]}
namesById={activeNames}
emptyMessage='Run this provider or the full batch to see results.'
/>
</TabsContent>
</Tabs>
</Card>
) : null}
</div>
)}
{selectedProviders.length > 0 && (
<div className='bg-background sticky bottom-0 flex shrink-0 flex-col gap-3 rounded-lg border p-3 shadow-sm sm:flex-row sm:items-center sm:justify-between'>
<div className='text-sm'>
<div className='font-medium'>
{selectedProviderIds.length} providers · {totalModels} models ·{' '}
{totalRoutes} routes
</div>
<div className='text-muted-foreground text-xs'>
{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.
</div>
</div>
<Button
onClick={runAllProviders}
disabled={!allReady || runningProviderIds.length > 0}
>
{runningProviderIds.length > 0 ? (
<Loader2 className='h-4 w-4 animate-spin' />
) : (
<Play className='h-4 w-4' />
)}
{runningProviderIds.length > 0
? `Running ${runningProviderIds.length} provider${runningProviderIds.length === 1 ? '' : 's'}`
: 'Run all providers'}
</Button>
</div>
)}
</div>
</AppPageShell>
);
}
+15 -6
View File
@@ -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={
<DialogTrigger asChild>
<Button>
<Plus className='h-4 w-4' />
Add Provider
<>
<Button asChild variant='outline'>
<Link href='/providers/certification'>
<BadgeCheck className='h-4 w-4' />
Certify Providers
</Link>
</Button>
</DialogTrigger>
<DialogTrigger asChild>
<Button>
<Plus className='h-4 w-4' />
Add Provider
</Button>
</DialogTrigger>
</>
}
/>
<ProviderFormDialogContent
+20
View File
@@ -25,9 +25,11 @@ import {
AlertTriangle,
Unlock,
Loader2,
ShieldCheck,
} from 'lucide-react';
import { ProviderBalance } from '@/components/provider-balance';
import { ProviderModelsPanel } from '@/components/provider-models-panel';
import { ProviderCertificationDialog } from '@/components/provider-certification-dialog';
import { RoutstrCreateKeySection } from '@/components/providers/RoutstrCreateKeySection';
import { RoutstrProviderService } from '@/lib/api/services/routstr-provider';
import { getErrorStatus } from '@/lib/api/client';
@@ -93,6 +95,7 @@ export function ProviderCard({
}: ProviderCardProps) {
const queryClient = useQueryClient();
const [isKeyModalOpen, setIsKeyModalOpen] = useState(false);
const [isCertifyOpen, setIsCertifyOpen] = useState(false);
const [isReleaseDialogOpen, setIsReleaseDialogOpen] = useState(false);
// The claim as the query cache held it when the admin opened the dialog.
// The mutation sends this token rather than re-reading the query at submit
@@ -310,6 +313,17 @@ export function ProviderCard({
)}
</Button>
<Button
variant='outline'
size='sm'
onClick={() => setIsCertifyOpen(true)}
className='justify-center gap-1.5'
title='Probe the upstream and verify usage, pricing, caching and margin'
>
<ShieldCheck className='h-4 w-4' />
<span>Certify</span>
</Button>
<Button
variant='outline'
size='sm'
@@ -369,6 +383,12 @@ export function ProviderCard({
</AlertDialogContent>
</AlertDialog>
<ProviderCertificationDialog
provider={provider}
open={isCertifyOpen}
onOpenChange={setIsCertifyOpen}
/>
<Dialog open={isKeyModalOpen} onOpenChange={setIsKeyModalOpen}>
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[500px]'>
<DialogHeader>
@@ -0,0 +1,180 @@
'use client';
import { useEffect, useState } from 'react';
import { useQuery } from '@tanstack/react-query';
import { Loader2, RotateCcw } from 'lucide-react';
import { ProviderCertificationResults } 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 {
Dialog,
DialogContent,
DialogDescription,
DialogHeader,
DialogTitle,
} from '@/components/ui/dialog';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import { useProviderCertificationRunner } from '@/hooks/use-provider-certification-runner';
import { AdminService } from '@/lib/api/services/admin';
import type { UpstreamProvider } from '@/lib/api/services/admin';
import {
buildModelRuns,
countCertificationTargets,
emptyCertificationSetup,
getModelsNeedingPath,
} from '@/lib/provider-certification';
import type { ProviderCertificationSetup } from '@/lib/provider-certification';
interface ProviderCertificationDialogProps {
provider: UpstreamProvider;
open: boolean;
onOpenChange: (open: boolean) => void;
}
export function ProviderCertificationDialog({
provider,
open,
onOpenChange,
}: ProviderCertificationDialogProps) {
const [workspaceTab, setWorkspaceTab] = useState<'setup' | 'results'>(
'setup'
);
const [setup, setSetup] = useState<ProviderCertificationSetup>(
emptyCertificationSetup
);
const { results, progress, isPending, run, reset } =
useProviderCertificationRunner(provider.id);
const models = useQuery({
queryKey: ['provider-models', provider.id],
queryFn: () => AdminService.getProviderModels(provider.id),
enabled: open,
});
useEffect(() => {
if (!open) {
reset();
setWorkspaceTab('setup');
setSetup(emptyCertificationSetup());
}
}, [open, reset]);
const modelRuns = buildModelRuns(setup, models.data);
const targetCount = countCertificationTargets(modelRuns);
const modelsNeedingPath = getModelsNeedingPath(setup);
const namesById = getCertificationModelNames(models.data);
const updateSetup = (nextSetup: ProviderCertificationSetup) => {
setSetup(nextSetup);
reset();
};
const runCertification = () => {
setWorkspaceTab('results');
void run(modelRuns, setup.checkCache);
};
return (
<Dialog
open={open}
onOpenChange={(nextOpen) => {
if (!nextOpen) reset();
onOpenChange(nextOpen);
}}
>
<DialogContent className='flex h-[90dvh] max-h-[90dvh] flex-col overflow-hidden sm:max-w-[780px]'>
<DialogHeader className='shrink-0'>
<DialogTitle>Certify upstream models</DialogTitle>
<DialogDescription>
Select models, then use the provider default, choose specific paths,
or test every path. Models run one at a time; paths for the same
model run in parallel against {provider.base_url}.
</DialogDescription>
</DialogHeader>
<Tabs
value={workspaceTab}
onValueChange={(value) =>
setWorkspaceTab(value as 'setup' | 'results')
}
className='min-h-0 flex-1 overflow-hidden'
>
<TabsList className='grid w-full shrink-0 grid-cols-2'>
<TabsTrigger value='setup'>Setup</TabsTrigger>
<TabsTrigger
value='results'
disabled={!isPending && results.length === 0}
>
Results
{(isPending || results.length > 0) && (
<Badge variant='secondary' className='ml-1 px-1.5 py-0 text-xs'>
{results.length}/{targetCount}
</Badge>
)}
</TabsTrigger>
</TabsList>
<TabsContent
value='setup'
className='mt-0 min-h-0 overflow-hidden data-[state=active]:flex data-[state=active]:flex-col'
>
<div className='min-h-0 flex-1 overflow-y-auto pr-1'>
<ProviderCertificationSetupPanel
models={models.data}
isLoading={models.isLoading}
error={models.error}
setup={setup}
onChange={updateSetup}
disabled={isPending}
idPrefix={`certify-${provider.id}`}
/>
</div>
<div className='mt-3 flex shrink-0 justify-end border-t pt-3'>
<Button
variant='outline'
size='sm'
onClick={runCertification}
disabled={
isPending ||
setup.selectedModelIds.length === 0 ||
modelsNeedingPath.length > 0
}
className='gap-1.5'
>
{isPending ? (
<Loader2 className='h-4 w-4 animate-spin' />
) : (
<RotateCcw className='h-4 w-4' />
)}
{isPending
? progress
? `Running ${progress.modelIndex} of ${progress.modelTotal}`
: 'Finishing in-flight probe'
: results.length > 0
? `Run ${targetCount} route${targetCount === 1 ? '' : 's'} again`
: `Certify ${targetCount || ''} route${targetCount === 1 ? '' : 's'}`}
</Button>
</div>
</TabsContent>
<TabsContent
value='results'
className='mt-0 min-h-0 overflow-hidden data-[state=active]:flex data-[state=active]:flex-col'
>
<ProviderCertificationResults
results={results}
progress={progress}
namesById={namesById}
/>
</TabsContent>
</Tabs>
</DialogContent>
</Dialog>
);
}
@@ -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 (
<Badge variant='outline' className={cn('gap-1', style.className)}>
<Icon className='h-3 w-3' />
{style.label}
</Badge>
);
}
function ResultStatusBadge({
status,
}: {
status: CertificationStatus | 'error';
}) {
if (status === 'error') {
return (
<Badge
variant='outline'
className='border-red-500/40 bg-red-500/10 text-red-700 dark:text-red-400'
>
Error
</Badge>
);
}
return <StatusBadge status={status} />;
}
function RowItem({ row }: { row: CertificationRow }) {
const [showEvidence, setShowEvidence] = useState(false);
const hasEvidence = Object.keys(row.evidence).length > 0;
return (
<li className='rounded-md border p-3 text-sm'>
<div className='flex items-start justify-between gap-3'>
<div className='min-w-0'>
<div className='font-medium'>{row.title}</div>
<div className='text-muted-foreground mt-0.5 break-words'>
{row.detail}
</div>
<div className='text-muted-foreground mt-1 font-mono text-xs'>
{row.id}
</div>
</div>
<StatusBadge status={row.status} />
</div>
{hasEvidence && (
<div className='mt-2'>
<button
type='button'
onClick={() => setShowEvidence((value) => !value)}
className='text-muted-foreground hover:text-foreground inline-flex items-center gap-1 text-xs'
>
{showEvidence ? (
<ChevronDown className='h-3 w-3' />
) : (
<ChevronRight className='h-3 w-3' />
)}
Evidence
</button>
{showEvidence && (
<pre className='bg-muted mt-2 max-h-64 overflow-auto rounded p-2 font-mono text-xs'>
{JSON.stringify(row.evidence, null, 2)}
</pre>
)}
</div>
)}
</li>
);
}
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 (
<div className='space-y-4'>
<ul className='grid gap-2 sm:grid-cols-2'>
{report.checklist.map((goal) => {
const style = STATUS_STYLES[goal.status];
const Icon = style.icon;
return (
<li
key={goal.goal}
className={cn(
'flex items-start gap-2 rounded-md border p-2 text-sm',
style.className
)}
>
<Icon className='mt-0.5 h-4 w-4 shrink-0' />
<span>{goal.label}</span>
</li>
);
})}
</ul>
<div className='text-muted-foreground flex flex-wrap items-center gap-x-3 gap-y-1 text-xs'>
<span>
{report.rows.length} checks · {failing} failed · {warning} warnings
</span>
<span>Generated {new Date(report.generated_at).toLocaleString()}</span>
</div>
<ul className='space-y-2'>
{report.rows.map((row) => (
<RowItem key={row.id} row={row} />
))}
</ul>
</div>
);
}
interface ProviderCertificationResultsProps {
results: ModelCertificationResult[];
progress?: CertificationProgress | null;
namesById: Map<string, string>;
emptyMessage?: string;
}
export function ProviderCertificationResults({
results,
progress,
namesById,
emptyMessage = 'Results will appear here as certification completes.',
}: ProviderCertificationResultsProps) {
const [activeResult, setActiveResult] = useState<string>('');
useEffect(() => {
if (
results.length > 0 &&
!results.some((result) => result.resultKey === activeResult)
) {
setActiveResult(results[0].resultKey);
}
}, [activeResult, results]);
return (
<div className='flex min-h-0 flex-1 flex-col gap-2 overflow-hidden'>
{progress && (
<div className='text-muted-foreground flex shrink-0 items-center gap-2 rounded-md border p-3 text-sm'>
<Loader2 className='h-4 w-4 animate-spin' />
<span className='min-w-0 truncate'>
Probing {namesById.get(progress.modelId) ?? progress.modelId}
{progress.pathCount > 1
? ` across ${progress.pathCount} paths in parallel`
: ''}{' '}
— model {progress.modelIndex} of {progress.modelTotal}
</span>
</div>
)}
{results.length === 0 ? (
<div className='text-muted-foreground flex min-h-40 flex-1 items-center justify-center text-center text-sm'>
{emptyMessage}
</div>
) : (
<Tabs
value={activeResult}
onValueChange={setActiveResult}
className='min-h-0 flex-1 overflow-hidden'
>
<div className='max-h-28 shrink-0 overflow-y-auto rounded-md border p-2'>
<TabsList className='flex h-auto w-full flex-wrap justify-start gap-1 bg-transparent p-0'>
{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 (
<TabsTrigger
key={result.resultKey}
value={result.resultKey}
title={namesById.get(result.modelId) ?? result.modelId}
className='h-7 max-w-44 min-w-0 gap-1 px-2 text-xs'
>
<Icon
className={cn(
status === 'ok' && 'text-emerald-600',
status === 'warn' && 'text-amber-600',
(status === 'fail' || status === 'error') &&
'text-red-600'
)}
/>
<span className='truncate'>
{namesById.get(result.modelId) ?? result.modelId}
</span>
{modelRouteCount > 1 && (
<span className='text-muted-foreground'>
{modelRouteNumber}
</span>
)}
</TabsTrigger>
);
})}
</TabsList>
</div>
<div className='min-h-0 flex-1 overflow-y-auto pr-1'>
{results.map((result) => {
const status = getCertificationResultStatus(result);
return (
<TabsContent
key={result.resultKey}
value={result.resultKey}
className='space-y-3'
>
<div className='bg-muted/30 space-y-2 rounded-md border p-3'>
<div className='flex flex-wrap items-center justify-between gap-2'>
<div className='font-medium'>
{namesById.get(result.modelId) ?? result.modelId}
</div>
<ResultStatusBadge status={status} />
</div>
<div className='grid gap-1 text-xs'>
<span className='text-muted-foreground'>Model path</span>
<div className='bg-background rounded px-2 py-1.5 font-medium break-words'>
{result.pathLabel}
</div>
</div>
</div>
{result.report ? (
<CertificationReport report={result.report} />
) : (
<div className='rounded-md border border-red-500/40 bg-red-500/10 p-3 text-sm text-red-700 dark:text-red-400'>
{result.error ?? 'Certification failed'}
</div>
)}
</TabsContent>
);
})}
</div>
</Tabs>
)}
</div>
);
}
@@ -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<string, string> {
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 (
<CommandGroup heading={`${label} (${options.length})`}>
{options.map(({ model }) => (
<CommandItem
key={`${label}-${model.id}`}
value={`${model.name} ${model.id}`}
onSelect={() => toggleModel(model.id)}
disabled={disabled}
>
<Checkbox
checked={setup.selectedModelIds.includes(model.id)}
tabIndex={-1}
aria-hidden='true'
className='pointer-events-none'
/>
<span className='min-w-0 flex-1 truncate'>
{model.name || model.id}
</span>
<span className='text-muted-foreground max-w-56 truncate font-mono text-xs'>
{model.id}
</span>
</CommandItem>
))}
</CommandGroup>
);
};
return (
<div className='space-y-4'>
<div className='space-y-2'>
<div className='flex items-center justify-between gap-3'>
<Label>Models</Label>
<div className='flex items-center gap-2'>
<span className='text-muted-foreground text-xs'>
{setup.selectedModelIds.length} selected
</span>
{setup.selectedModelIds.length > 0 && !disabled && (
<Button
type='button'
variant='ghost'
size='sm'
onClick={() =>
onChange({
...emptyCertificationSetup(),
checkCache: setup.checkCache,
})
}
>
Clear
</Button>
)}
</div>
</div>
<Command className='h-auto rounded-md border'>
<CommandInput
placeholder='Search models by name or ID…'
disabled={isLoading || disabled}
/>
<CommandList className='max-h-64'>
<CommandEmpty>
{isLoading ? 'Loading models…' : 'No models found'}
</CommandEmpty>
{renderModelGroup('Configured models', configuredOptions)}
{renderModelGroup('Discovered models', discoveredOptions)}
</CommandList>
</Command>
{Boolean(error) && (
<p className='text-destructive text-sm'>{getErrorMessage(error)}</p>
)}
</div>
{setup.selectedModelIds.map((modelId) => {
const paths = getExactCertificationPaths(models, modelId);
const mode = setup.pathModes[modelId] ?? 'default';
const selectedPaths = setup.selectedModelPaths[modelId] ?? [];
return (
<div key={modelId} className='space-y-3 rounded-md border p-3'>
<div className='min-w-0'>
<div className='truncate text-sm font-medium'>
{namesById.get(modelId) ?? modelId}
</div>
<div className='text-muted-foreground truncate font-mono text-xs'>
{modelId}
</div>
</div>
{paths.length > 0 ? (
<>
<ToggleGroup
type='single'
variant='outline'
size='sm'
value={mode}
onValueChange={(value) => {
if (!value) return;
onChange({
...setup,
pathModes: {
...setup.pathModes,
[modelId]: value as ModelPathMode,
},
});
}}
disabled={disabled}
className='w-full justify-start'
>
<ToggleGroupItem value='default'>Default</ToggleGroupItem>
<ToggleGroupItem value='selected'>
Choose paths
</ToggleGroupItem>
<ToggleGroupItem value='all'>All paths</ToggleGroupItem>
</ToggleGroup>
{mode === 'default' && (
<p className='text-muted-foreground text-xs'>
Uses the upstream provider&apos;s normal model routing.
</p>
)}
{mode === 'selected' && (
<Popover>
<PopoverTrigger asChild>
<Button
type='button'
variant='outline'
size='sm'
className='w-full justify-between'
disabled={disabled}
>
<span className='truncate'>
{selectedPaths.length === 0
? 'Select paths'
: `${selectedPaths.length} path${selectedPaths.length === 1 ? '' : 's'} selected`}
</span>
<ChevronDown className='h-4 w-4' />
</Button>
</PopoverTrigger>
<PopoverContent
align='start'
className='w-80 max-w-[calc(100vw-2rem)] p-2'
>
<div className='max-h-64 space-y-1 overflow-y-auto overscroll-contain'>
{paths.map((path) => {
const checked = selectedPaths.includes(path.path);
return (
<label
key={path.path}
className='hover:bg-muted flex cursor-pointer items-center gap-2 rounded-sm px-2 py-1.5 text-sm'
>
<Checkbox
checked={checked}
onCheckedChange={(value) => {
onChange({
...setup,
selectedModelPaths: {
...setup.selectedModelPaths,
[modelId]:
value === true
? [...selectedPaths, path.path]
: selectedPaths.filter(
(item) => item !== path.path
),
},
});
}}
/>
<span className='min-w-0 truncate'>
{getCertificationPathLabel(path)}
</span>
</label>
);
})}
</div>
</PopoverContent>
</Popover>
)}
{mode === 'all' && (
<p className='text-muted-foreground text-xs'>
All {paths.length} paths will run one after another, each
with its own probe calls.
</p>
)}
</>
) : (
<p className='text-muted-foreground text-xs'>
Only the provider default route is available.
</p>
)}
</div>
);
})}
{modelsNeedingPath.length > 0 && (
<p className='text-muted-foreground text-xs'>
Choose at least one path for each model using “Choose paths”.
</p>
)}
<div className='flex items-center gap-2'>
<Checkbox
id={`${idPrefix}-cache`}
checked={setup.checkCache}
onCheckedChange={(value) =>
onChange({ ...setup, checkCache: value === true })
}
disabled={disabled}
/>
<Label htmlFor={`${idPrefix}-cache`} className='text-sm'>
Probe prompt caching and margin
</Label>
</div>
{setup.checkCache && (
<p className='text-muted-foreground text-xs'>
Sends 2–3 extra completions with a ~4.4k-token prompt per model path,
billed by the upstream.
</p>
)}
</div>
);
}
@@ -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]);
});
@@ -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<ModelCertificationResult[]> {
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<ModelCertificationResult[]>([]);
const [progress, setProgress] = useState<CertificationProgress | null>(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 };
}
+55
View File
@@ -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<typeof CertificationStatusSchema>;
export type CertificationRow = z.infer<typeof CertificationRowSchema>;
export type CertificationGoal = z.infer<typeof CertificationGoalSchema>;
export type ProviderCertification = z.infer<typeof ProviderCertificationSchema>;
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<typeof ProviderTypeSchema>;
@@ -107,6 +150,7 @@ export type AdminModelPricing = z.infer<typeof AdminModelPricingSchema>;
export type AdminModelArchitecture = z.infer<
typeof AdminModelArchitectureSchema
>;
export type CertificationPath = z.infer<typeof CertificationPathSchema>;
export type ProviderModels = z.infer<typeof ProviderModelsSchema>;
export interface AdminModelAsModel {
@@ -317,6 +361,17 @@ export class AdminService {
);
}
static async certifyProvider(
providerId: number,
body: CertifyProviderRequest = {}
): Promise<ProviderCertification> {
const data = await apiClient.post<unknown>(
`/admin/api/upstream-providers/${providerId}/certify`,
body
);
return ProviderCertificationSchema.parse(data);
}
static async getProviderModels(providerId: number): Promise<ProviderModels> {
const data = await apiClient.get<ProviderModels>(
`/admin/api/upstream-providers/${providerId}/models`
+123
View File
@@ -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<string, ModelPathMode>;
selectedModelPaths: Record<string, string[]>;
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<number, ModelCertificationResult[]>
): 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');
Generated
+14
View File
@@ -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" },
]