mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge pull request #760 from 9qeklajc/feat/upstream-certification-harness
feat(certification): comprehensive upstream certification harness with live probes
This commit is contained in:
@@ -38,6 +38,7 @@ dev = [
|
||||
"psutil>=5.9.0",
|
||||
"aiohttp>=3.9.0",
|
||||
"pytest-benchmark>=4.0.0",
|
||||
"respx>=0.21",
|
||||
"routstr",
|
||||
]
|
||||
|
||||
|
||||
+513
-1
@@ -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
@@ -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
@@ -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
@@ -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, {}),
|
||||
]
|
||||
@@ -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={
|
||||
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
@@ -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 == [""]
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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'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 };
|
||||
}
|
||||
@@ -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`
|
||||
|
||||
@@ -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');
|
||||
@@ -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" },
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user