feat: add model path certification

This commit is contained in:
9qeklajc
2026-09-23 21:19:36 +02:00
parent 0b8e07f834
commit 721eb4ff8d
11 changed files with 2489 additions and 8 deletions
+93 -1
View File
@@ -28,6 +28,7 @@ from .db import (
CashuTransaction,
CliToken,
LightningInvoice,
ModelPathRow,
ModelRow,
UpstreamProviderRow,
create_session,
@@ -1234,6 +1235,30 @@ async def get_provider_models(provider_id: str) -> dict[str, object]:
m for m in upstream_models if m.id not in db_model_ids
]
path_result = await session.exec(
select(ModelPathRow).where(
ModelPathRow.upstream_provider_id == provider_pk
)
)
path_rows = list(path_result.all())
paths_by_public_id: dict[str, list[dict[str, object]]] = {}
for row in path_rows:
paths_by_public_id.setdefault(row.model_id.lower(), []).append(
{
"path": row.path,
"endpoint_tag": row.endpoint_tag,
"endpoint_name": row.endpoint_name,
}
)
from ..upstream.model_paths import public_model_id
certification_paths: dict[str, list[dict[str, object]]] = {}
for model in [*db_models, *filtered_remote_models]:
forwarded_id = model.forwarded_model_id or model.id
paths = paths_by_public_id.get(public_model_id(forwarded_id).lower(), [])
certification_paths[model.id] = paths
return {
"provider": {
"id": provider.id,
@@ -1247,6 +1272,7 @@ async def get_provider_models(provider_id: str) -> dict[str, object]:
# missing one; show the operator the value that needs fixing.
"db_models": [json_compliant(m.dict()) for m in db_models],
"remote_models": [json_compliant(m.dict()) for m in filtered_remote_models],
"certification_paths": certification_paths,
}
@@ -1499,7 +1525,9 @@ async def get_upstream_provider_report(provider_id: str) -> dict[str, object]:
class CertifyRequest(BaseModel):
model_id: str | None = None
model_path: str | None = None
timeout_seconds: float | None = None
check_cache: bool = True
@admin_router.post(
@@ -1540,6 +1568,36 @@ async def certify_upstream_provider(
)
enabled_rows = list(result.all())
endpoint_tag: str | None = None
selected_path: ModelPathRow | None = None
if payload.model_path is not None:
from ..proxy import _model_ids_match
from ..upstream.model_paths import decode_model_path
selector = decode_model_path(payload.model_path)
if selector is None:
raise HTTPException(status_code=400, detail="Malformed model path")
if payload.model_id is None or not _model_ids_match(
payload.model_id, selector.model_id
):
raise HTTPException(
status_code=400,
detail="Model path does not match the selected model",
)
path_result = await session.exec(
select(ModelPathRow).where(
ModelPathRow.upstream_provider_id == provider_pk,
ModelPathRow.path == payload.model_path,
)
)
selected_path = path_result.first()
if selected_path is None:
raise HTTPException(
status_code=400,
detail="Model path is not available for this provider",
)
endpoint_tag = selector.endpoint_tag
evaluations = [
_evaluate_model_row(row, provider, provider_pk) for row in enabled_rows
]
@@ -1561,7 +1619,7 @@ async def certify_upstream_provider(
if model_id is None:
model_id = enabled_rows[0].id
from ..proxy import get_candidates
from ..proxy import get_candidates, get_upstreams
model_obj = None
if model_id:
@@ -1581,11 +1639,33 @@ async def certify_upstream_provider(
if model.upstream_provider_id == provider_pk:
model_obj = model
break
# A model selected from the provider's discovered catalog may not have
# a database override and therefore may not appear in get_candidates().
# The active upstream cache carries the same fee-adjusted USD and sats
# pricing used by the proxy, so it is the authoritative fallback for a
# pre-configuration certification probe.
if model_obj is None:
for upstream in get_upstreams():
if getattr(upstream, "db_id", None) != provider_pk:
continue
model_obj = next(
(
model
for model in upstream.get_cached_models()
if model.id == model_id
or model.forwarded_model_id == model_id
),
None,
)
if model_obj is not None:
break
if model_obj is None:
from ..upstream.certification import (
STATUS_WARN,
certification_row,
)
from ..upstream.certification_cache import skipped_cache_rows
live_rows = [
certification_row(
@@ -1624,9 +1704,19 @@ async def certify_upstream_provider(
"Skipped — no model to probe.",
{},
),
*skipped_cache_rows("Skipped — no model to probe."),
]
else:
sats_to_usd = sats_usd_price()
if selected_path is not None:
from ..upstream.model_paths import apply_model_path_pricing
model_obj = apply_model_path_pricing(
model_obj,
selected_path,
provider.provider_fee,
sats_to_usd,
)
# Clamp the admin-supplied timeout so a probe cannot hold the request
# open indefinitely.
requested = (
@@ -1642,6 +1732,8 @@ async def certify_upstream_provider(
provider_fee=provider.provider_fee,
sats_to_usd=sats_to_usd,
timeout=timeout,
check_cache=payload.check_cache,
endpoint_tag=endpoint_tag,
)
rows = pricing_rows + live_rows
+69 -1
View File
@@ -28,6 +28,7 @@ from ..core.logging import get_logger
from ..payment.cost_calculation import calculate_cost
from ..payment.rates import coerce_rate
from ..payment.usage import normalize_usage
from .model_paths import is_openrouter_base_url
if TYPE_CHECKING:
from ..payment.models import Model
@@ -124,6 +125,16 @@ CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = (
"Pricing in /v1/models — cost updates reflected in the models list",
("pricing.served_matches_configured", "pricing.enabled_models_served"),
),
(
"caching",
"Prompt caching — cache hits reported and billed at the cache rate",
("cache.reported", "cache.billing"),
),
(
"margin",
"Margin — node charge covers the upstream's cost",
("cost.margin",),
),
)
@@ -159,6 +170,7 @@ class ProbeResult:
base_url: str
models_url: str
chat_url: str
endpoint_tag: str | None = None
models_status: int | None = None
models_payload: dict[str, Any] | None = None
models_error: str | None = None
@@ -174,6 +186,7 @@ async def probe_upstream(
api_key: str,
model_id: str,
*,
endpoint_tag: str | None = None,
client: httpx.AsyncClient | None = None,
timeout: float = PROBE_TIMEOUT_SECONDS,
) -> ProbeResult:
@@ -187,6 +200,7 @@ async def probe_upstream(
base_url=base_url,
models_url=f"{base}/models",
chat_url=f"{base}/chat/completions",
endpoint_tag=endpoint_tag,
)
headers = {"Content-Type": "application/json"}
if api_key:
@@ -226,6 +240,11 @@ async def probe_upstream(
"max_tokens": PROBE_MAX_TOKENS,
"stream": False,
}
if endpoint_tag:
request_body["provider"] = {
"order": [endpoint_tag],
"allow_fallbacks": False,
}
try:
response = await client.post(
result.chat_url, json=request_body, headers=headers
@@ -702,12 +721,19 @@ async def run_live_checks(
client: httpx.AsyncClient | None = None,
timeout: float = PROBE_TIMEOUT_SECONDS,
pricing_known: bool = True,
check_cache: bool = True,
endpoint_tag: str | None = None,
) -> list[dict[str, Any]]:
"""Probe one upstream once and build the five live/derived rows."""
"""Probe one upstream and build the live/derived rows.
``check_cache`` adds the prompt-cache and margin rows, which cost two
more completions against a long prompt.
"""
probe = await probe_upstream(
base_url,
api_key,
model.forwarded_model_id or model.id,
endpoint_tag=endpoint_tag,
client=client,
timeout=timeout,
)
@@ -769,6 +795,37 @@ async def run_live_checks(
),
)
)
from .certification_cache import run_cache_checks, skipped_cache_rows
if not check_cache:
rows.extend(skipped_cache_rows("Skipped — cache checks disabled."))
elif probe.chat_payload is None:
rows.extend(
skipped_cache_rows("Skipped — the completion probe did not succeed.")
)
elif is_openrouter_base_url(base_url) and endpoint_tag is None:
rows.extend(
skipped_cache_rows(
"Skipped — select an exact OpenRouter model path so both cache "
"requests use the same upstream endpoint."
)
)
else:
rows.extend(
await run_cache_checks(
base_url,
api_key,
model,
provider_fee=provider_fee,
sats_to_usd=sats_to_usd,
probe_payload=probe.chat_payload,
client=client,
timeout=timeout,
pricing_known=pricing_known,
endpoint_tag=endpoint_tag,
)
)
return rows
@@ -877,6 +934,7 @@ async def certify_upstream_url(
timeout: float = PROBE_TIMEOUT_SECONDS,
sats_usd_price: float | None = None,
client: httpx.AsyncClient | None = None,
check_cache: bool = True,
) -> dict[str, Any]:
"""Certify an arbitrary upstream URL without touching the node's DB."""
from ..payment.models import litellm_cost_entry
@@ -911,6 +969,9 @@ async def certify_upstream_url(
{},
),
]
from .certification_cache import skipped_cache_rows
rows.extend(skipped_cache_rows("Skipped — no model to probe."))
return {
"target": target,
"rows": rows,
@@ -952,6 +1013,7 @@ async def certify_upstream_url(
client=client,
timeout=timeout,
pricing_known=pricing_known,
check_cache=check_cache,
)
return {"target": target, "rows": rows, "checklist": build_checklist(rows)}
@@ -1061,6 +1123,11 @@ def main(argv: list[str] | None = None) -> int:
"parseable — use this in pipelines."
),
)
parser.add_argument(
"--no-cache",
action="store_true",
help="Skip the prompt-cache and margin rows (saves two long completions)",
)
args = parser.parse_args(argv)
async def _run_all() -> list[dict[str, Any]]:
@@ -1076,6 +1143,7 @@ def main(argv: list[str] | None = None) -> int:
provider_fee=args.provider_fee,
timeout=args.timeout,
sats_usd_price=args.sats_usd_price,
check_cache=not args.no_cache,
)
)
return results
+586
View File
@@ -0,0 +1,586 @@
"""Prompt-cache and margin certification for an upstream provider.
Three questions the one-token probe cannot answer:
* does the upstream *report* prompt-cache hits in a dialect the node parses,
* does the node bill cached reads at the discounted rate (client side), and
* does the node's charge cover what the upstream charged (node side).
The cache probe sends the same long system prompt twice; the second call is
the one expected to report cached reads. Calls go straight to the upstream
with ``httpx`` and never enter the billing path.
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
import httpx
from ..core.logging import get_logger
from ..payment.cost_calculation import CostDataError, calculate_cost
from ..payment.usage import NormalizedUsage, normalize_usage
from .certification import (
_PROBE_MAX_COST_MSATS,
COST_TOLERANCE_MSATS,
PROBE_MAX_TOKENS,
PROBE_TIMEOUT_SECONDS,
STATUS_FAIL,
STATUS_OK,
STATUS_WARN,
_expected_token_msats,
_expected_usd_msats,
_reported_usd_cost,
certification_row,
safe_row,
)
if TYPE_CHECKING:
from ..payment.models import Model
logger = get_logger(__name__)
# OpenAI caches prefixes of 1024+ tokens; Anthropic Haiku needs 2048+. The
# filler lands around 3000 tokens so every dialect can hit its threshold.
CACHE_PROBE_LINES = 220
CACHE_PROBE_QUESTION = "Reply with the single word: ok"
ROW_REPORTED = "cache.reported"
ROW_BILLING = "cache.billing"
ROW_MARGIN = "cost.margin"
TITLE_REPORTED = "Upstream reports prompt-cache hits"
TITLE_BILLING = "Cached tokens billed at the cache-read rate"
TITLE_MARGIN = "Node charge covers upstream cost"
def cache_probe_prefix() -> str:
lines = [
"You are a certification probe. Ignore the reference table below and "
"answer the final question with one word."
]
for index in range(CACHE_PROBE_LINES):
lines.append(
f"Reference row {index:04d}: token {index * 7919 % 10007} maps to "
f"slot {index * 104729 % 1009} in region {index % 17}."
)
return "\n".join(lines)
@dataclass
class CacheProbeResult:
"""Two identical completions; the second should read from the cache."""
chat_url: str
request_format: str = "cache_control"
endpoint_tag: str | None = None
statuses: list[int | None] = field(default_factory=list)
payloads: list[dict[str, Any] | None] = field(default_factory=list)
errors: list[str | None] = field(default_factory=list)
latencies_ms: list[float | None] = field(default_factory=list)
@property
def second_payload(self) -> dict[str, Any] | None:
return self.payloads[1] if len(self.payloads) > 1 else None
@property
def second_error(self) -> str | None:
if len(self.errors) > 1:
return self.errors[1]
return self.errors[0] if self.errors else "cache probe did not run"
def _request_body(
model_id: str, prefix: str, fmt: str, endpoint_tag: str | None
) -> dict[str, Any]:
system: Any
if fmt == "cache_control":
system = [
{
"type": "text",
"text": prefix,
"cache_control": {"type": "ephemeral"},
}
]
else:
system = prefix
body: dict[str, Any] = {
"model": model_id,
"messages": [
{"role": "system", "content": system},
{"role": "user", "content": CACHE_PROBE_QUESTION},
],
"max_tokens": PROBE_MAX_TOKENS,
"stream": False,
}
if endpoint_tag:
body["provider"] = {
"order": [endpoint_tag],
"allow_fallbacks": False,
}
return body
async def _post_completion(
client: httpx.AsyncClient,
url: str,
body: dict[str, Any],
headers: dict[str, str],
) -> tuple[int | None, dict[str, Any] | None, str | None, float]:
started = time.monotonic()
try:
response = await client.post(url, json=body, headers=headers)
except Exception as exc: # noqa: BLE001 - transport failure is a row status
latency = round((time.monotonic() - started) * 1000, 2)
return None, None, f"{type(exc).__name__}: {exc}", latency
latency = round((time.monotonic() - started) * 1000, 2)
try:
payload = response.json()
except Exception as exc: # noqa: BLE001 - any decode failure is the signal
return response.status_code, None, f"{type(exc).__name__}: {exc}", latency
if not isinstance(payload, dict):
return (
response.status_code,
None,
f"expected a JSON object, got {type(payload).__name__}",
latency,
)
return response.status_code, payload, None, latency
def _record(
result: CacheProbeResult,
outcome: tuple[int | None, dict[str, Any] | None, str | None, float],
) -> None:
status, payload, error, latency = outcome
result.statuses.append(status)
result.payloads.append(payload)
result.errors.append(error)
result.latencies_ms.append(latency)
def _is_2xx(status: int | None) -> bool:
return status is not None and 200 <= status < 300
async def probe_cache(
base_url: str,
api_key: str,
model_id: str,
*,
endpoint_tag: str | None = None,
client: httpx.AsyncClient | None = None,
timeout: float = PROBE_TIMEOUT_SECONDS,
) -> CacheProbeResult:
"""Send the same long prompt twice.
The first attempt marks the prefix with an Anthropic-style
``cache_control`` part. Upstreams that reject the part get a plain string
retry, and the second call mirrors whichever format succeeded.
"""
base = base_url.rstrip("/")
result = CacheProbeResult(
chat_url=f"{base}/chat/completions", endpoint_tag=endpoint_tag
)
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
prefix = cache_probe_prefix()
owns_client = client is None
if client is None:
client = httpx.AsyncClient(timeout=timeout)
try:
first = await _post_completion(
client,
result.chat_url,
_request_body(model_id, prefix, "cache_control", endpoint_tag),
headers,
)
if not _is_2xx(first[0]) and first[0] is not None:
result.request_format = "plain"
first = await _post_completion(
client,
result.chat_url,
_request_body(model_id, prefix, "plain", endpoint_tag),
headers,
)
_record(result, first)
if not _is_2xx(first[0]):
return result
second = await _post_completion(
client,
result.chat_url,
_request_body(model_id, prefix, result.request_format, endpoint_tag),
headers,
)
_record(result, second)
finally:
if owns_client:
await client.aclose()
return result
def _raw_cache_keys(value: Any, path: str = "") -> list[str]:
"""Paths of positive numeric fields whose name mentions a cache."""
found: list[str] = []
if isinstance(value, dict):
for key, item in value.items():
child = f"{path}.{key}" if path else str(key)
if "cach" in str(key).lower() and isinstance(item, (int, float)):
if not isinstance(item, bool) and item > 0:
found.append(child)
found.extend(_raw_cache_keys(item, child))
return found
def _usage_of(payload: dict[str, Any] | None) -> NormalizedUsage | None:
if not isinstance(payload, dict):
return None
try:
return normalize_usage(payload.get("usage"))
except Exception: # noqa: BLE001 - a malformed usage object is a row status
return None
def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]:
evidence: dict[str, Any] = {
"url": probe.chat_url,
"request_format": probe.request_format,
"endpoint_tag": probe.endpoint_tag,
"statuses": probe.statuses,
"latencies_ms": probe.latencies_ms,
}
payload = probe.second_payload
if payload is None or not _is_2xx(probe.statuses[-1] if probe.statuses else None):
evidence["error"] = probe.second_error
return certification_row(
ROW_REPORTED,
STATUS_FAIL,
TITLE_REPORTED,
f"The cache probe did not get two successful completions: "
f"{probe.second_error or 'no response body'}.",
evidence,
)
first_usage = _usage_of(probe.payloads[0])
second_usage = _usage_of(payload)
evidence["first_usage"] = first_usage.dict() if first_usage else None
evidence["second_usage"] = second_usage.dict() if second_usage else None
if second_usage is not None and second_usage.cache_read_tokens > 0:
return certification_row(
ROW_REPORTED,
STATUS_OK,
TITLE_REPORTED,
f"The repeated prompt reported {second_usage.cache_read_tokens} "
f"cached tokens (first call wrote {first_usage.cache_write_tokens if first_usage else 0}).",
evidence,
)
raw_keys = _raw_cache_keys(payload.get("usage"))
if raw_keys:
evidence["unrecognised_cache_fields"] = raw_keys
return certification_row(
ROW_REPORTED,
STATUS_FAIL,
TITLE_REPORTED,
"The upstream reported cache tokens under fields the node does not "
f"parse ({', '.join(raw_keys)}); cached reads would be billed at "
"the full input rate.",
evidence,
)
return certification_row(
ROW_REPORTED,
STATUS_WARN,
TITLE_REPORTED,
"Two identical prompts produced no cache hit. Either the model does "
"not support prompt caching or the upstream hides it; clients pay the "
"full input rate on repeated prompts.",
evidence,
)
def cache_billing_row(
*,
model: Model,
probe: CacheProbeResult,
cost_data: Any,
pricing_known: bool = True,
) -> dict[str, Any]:
usage = _usage_of(probe.second_payload)
evidence: dict[str, Any] = {"model_id": model.id}
if usage is None or usage.cache_read_tokens <= 0:
return certification_row(
ROW_BILLING,
STATUS_WARN,
TITLE_BILLING,
"No cached reads were reported, so there is nothing to price.",
evidence,
)
if model.sats_pricing is None or not pricing_known:
return certification_row(
ROW_BILLING,
STATUS_WARN,
TITLE_BILLING,
"No pricing is known for this model, so the cache discount cannot "
"be verified.",
evidence,
)
if isinstance(cost_data, CostDataError):
evidence["error"] = cost_data.message
return certification_row(
ROW_BILLING,
STATUS_FAIL,
TITLE_BILLING,
f"The cost engine could not price the cached completion: "
f"{cost_data.message}.",
evidence,
)
pricing = model.sats_pricing
cache_read_rate = float(pricing.input_cache_read or 0.0)
input_rate = float(pricing.prompt)
full_usage = NormalizedUsage(
input_tokens=usage.input_tokens
+ usage.cache_read_tokens
+ usage.cache_write_tokens,
output_tokens=usage.output_tokens,
)
try:
expected_total, _, _ = _expected_token_msats(pricing, usage)
full_total, _, _ = _expected_token_msats(pricing, full_usage)
except (ValueError, OverflowError) as exc:
evidence["error"] = f"{type(exc).__name__}: {exc}"
return certification_row(
ROW_BILLING,
STATUS_FAIL,
TITLE_BILLING,
f"The expected charge could not be derived: {exc}.",
evidence,
)
actual_total = int(cost_data.total_msats)
reported_usd = _reported_usd_cost(probe.second_payload or {})
evidence.update(
{
"usage": usage.dict(),
"cache_read_rate_sats": cache_read_rate,
"input_rate_sats": input_rate,
"actual_total_msats": actual_total,
"expected_total_msats": expected_total,
"full_price_total_msats": full_total,
"reported_usd": reported_usd or None,
}
)
if reported_usd > 0:
return certification_row(
ROW_BILLING,
STATUS_OK,
TITLE_BILLING,
f"Billed {actual_total} msats from the upstream-reported cost, "
f"which already carries the cache discount "
f"(full token price would be {full_total} msats).",
evidence,
)
if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS:
return certification_row(
ROW_BILLING,
STATUS_FAIL,
TITLE_BILLING,
f"The engine charged {actual_total} msats but the configured cache "
f"rate implies {expected_total} msats.",
evidence,
)
if cache_read_rate <= 0.0 or cache_read_rate >= input_rate:
return certification_row(
ROW_BILLING,
STATUS_WARN,
TITLE_BILLING,
f"Cached reads are billed at the full input rate ({actual_total} "
"msats) because no discounted cache-read rate is configured; "
"clients pay more than the upstream charges.",
evidence,
)
return certification_row(
ROW_BILLING,
STATUS_OK,
TITLE_BILLING,
f"Charged {actual_total} msats for {usage.cache_read_tokens} cached "
f"tokens, {full_total - actual_total} msats below the full input price.",
evidence,
)
def cost_margin_row(
*,
model: Model,
payloads: list[dict[str, Any] | None],
provider_fee: float,
sats_to_usd: float,
pricing_known: bool = True,
) -> dict[str, Any]:
"""Configured token pricing must cover what the upstream reports charging.
Responses that carry a USD cost are billed from it, so they cannot lose
money themselves; they are used here as a price sample. The configured
token pricing is what every other path bills from (streams, upstreams
that omit cost, the served ``/v1/models`` list), so a sample where it
falls below the fee-adjusted upstream cost means those paths underprice.
Upstreams that report no cost give no sample and the row stays a warn.
"""
evidence: dict[str, Any] = {
"model_id": model.id,
"provider_fee": provider_fee,
"sats_usd_price": sats_to_usd,
"samples": [],
}
if model.sats_pricing is None or not pricing_known:
return certification_row(
ROW_MARGIN,
STATUS_WARN,
TITLE_MARGIN,
"No pricing is known for this model, so the margin cannot be verified.",
evidence,
)
samples: list[dict[str, Any]] = []
short: list[str] = []
for payload in payloads:
if not isinstance(payload, dict):
continue
reported_usd = _reported_usd_cost(payload)
usage = _usage_of(payload)
if reported_usd <= 0 or usage is None:
continue
try:
configured_total, _, _ = _expected_token_msats(model.sats_pricing, usage)
upstream_total = _expected_usd_msats(
reported_usd, provider_fee, sats_to_usd
)
except (ValueError, OverflowError) as exc:
evidence["error"] = f"{type(exc).__name__}: {exc}"
return certification_row(
ROW_MARGIN,
STATUS_FAIL,
TITLE_MARGIN,
f"The margin could not be derived: {exc}.",
evidence,
)
samples.append(
{
"usage": usage.dict(),
"reported_usd": reported_usd,
"upstream_msats_with_fee": upstream_total,
"configured_msats": configured_total,
}
)
if configured_total + COST_TOLERANCE_MSATS < upstream_total:
short.append(f"{configured_total} < {upstream_total}")
evidence["samples"] = samples
if not samples:
return certification_row(
ROW_MARGIN,
STATUS_WARN,
TITLE_MARGIN,
"The upstream does not report a cost, so the margin cannot be "
"verified live. Keep configured prices at or above the upstream's "
"list price.",
evidence,
)
if short:
return certification_row(
ROW_MARGIN,
STATUS_FAIL,
TITLE_MARGIN,
"Configured pricing is below the upstream's reported cost "
f"(configured < upstream msats: {'; '.join(short)}); token-billed "
"requests lose money.",
evidence,
)
return certification_row(
ROW_MARGIN,
STATUS_OK,
TITLE_MARGIN,
f"Configured pricing covers the upstream's reported cost on "
f"{len(samples)} sampled completion(s).",
evidence,
)
async def _price_payload(
payload: dict[str, Any] | None, model: Model, provider_fee: float
) -> Any:
if payload is None:
return CostDataError(
message="the cache probe did not succeed", code="no_completion"
)
try:
return await calculate_cost(
payload, _PROBE_MAX_COST_MSATS, model_obj=model, provider_fee=provider_fee
)
except Exception as exc: # noqa: BLE001 - a raising engine is a fail row
return CostDataError(
message=f"{type(exc).__name__}: {exc}", code="pricing_error"
)
async def run_cache_checks(
base_url: str,
api_key: str,
model: Model,
*,
provider_fee: float,
sats_to_usd: float,
probe_payload: dict[str, Any] | None,
client: httpx.AsyncClient | None = None,
timeout: float = PROBE_TIMEOUT_SECONDS,
pricing_known: bool = True,
endpoint_tag: str | None = None,
) -> list[dict[str, Any]]:
"""Run the cache probe and build the three cache/margin rows."""
probe = await probe_cache(
base_url,
api_key,
model.forwarded_model_id or model.id,
endpoint_tag=endpoint_tag,
client=client,
timeout=timeout,
)
cost_data = await _price_payload(probe.second_payload, model, provider_fee)
return [
safe_row(ROW_REPORTED, TITLE_REPORTED, lambda: cache_reported_row(probe)),
safe_row(
ROW_BILLING,
TITLE_BILLING,
lambda: cache_billing_row(
model=model,
probe=probe,
cost_data=cost_data,
pricing_known=pricing_known,
),
),
safe_row(
ROW_MARGIN,
TITLE_MARGIN,
lambda: cost_margin_row(
model=model,
payloads=[probe_payload, *probe.payloads],
provider_fee=provider_fee,
sats_to_usd=sats_to_usd,
pricing_known=pricing_known,
),
),
]
def skipped_cache_rows(reason: str) -> list[dict[str, Any]]:
return [
certification_row(ROW_REPORTED, STATUS_WARN, TITLE_REPORTED, reason, {}),
certification_row(ROW_BILLING, STATUS_WARN, TITLE_BILLING, reason, {}),
certification_row(ROW_MARGIN, STATUS_WARN, TITLE_MARGIN, reason, {}),
]
+52
View File
@@ -37,6 +37,7 @@ from ..core.logging import get_logger
if TYPE_CHECKING:
from sqlmodel.ext.asyncio.session import AsyncSession
from ..payment.models import Model
from .base import BaseUpstreamProvider
logger = get_logger(__name__)
@@ -884,6 +885,57 @@ def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None:
model["sats_pricing"] = priced.sats_pricing.dict()
def apply_model_path_pricing(
model: "Model",
row: ModelPathRow,
provider_fee: float,
sats_to_usd: float,
) -> "Model":
"""Return ``model`` priced from an exact endpoint path's own rates.
Direct paths already use the provider model cache and therefore carry the
same pricing as ``model``. OpenRouter endpoint rows instead contain raw,
endpoint-specific USD rates; certification must use those rates when its
requests are pinned to that endpoint.
"""
if row.endpoint_tag is None:
return model
from ..payment.models import (
Pricing,
_calculate_usd_max_costs,
_update_model_sats_pricing,
backfill_cache_pricing,
)
try:
metadata = json.loads(row.model_metadata)
if not isinstance(metadata, dict) or not isinstance(
metadata.get("pricing"), dict
):
return model
pricing = backfill_cache_pricing(
model.forwarded_model_id or row.model_id,
Pricing.parse_obj(metadata["pricing"]),
)
pricing = Pricing.parse_obj(
{key: float(value) * provider_fee for key, value in pricing.dict().items()}
)
priced = model.copy(update={"pricing": pricing, "sats_pricing": None})
(
pricing.max_prompt_cost,
pricing.max_completion_cost,
pricing.max_cost,
) = _calculate_usd_max_costs(priced)
return _update_model_sats_pricing(priced, sats_to_usd)
except Exception as exc:
logger.warning(
"Could not apply model-path pricing for certification",
extra={"model_id": model.id, "path": row.path, "error": str(exc)},
)
return model
def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]:
endpoint = None
if row.endpoint_tag or row.endpoint_name:
+408 -3
View File
@@ -20,8 +20,9 @@ from httpx import AsyncClient, Response
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.admin import admin_sessions
from routstr.core.db import ModelRow, UpstreamProviderRow
from routstr.core.db import ModelPathRow, ModelRow, UpstreamProviderRow
from routstr.proxy import reinitialize_upstreams
from routstr.upstream.model_paths import encode_model_path
# The conftest patches ``routstr.payment.price.sats_usd_price``, but
@@ -174,6 +175,33 @@ def _mock_chat_response(
}
def _caching_upstream(*, cached_tokens: int = 2900, report_cost: bool = True) -> Any:
"""Side effect that answers the one-token probe and then two long
prompts, reporting a cache hit (and optionally a USD cost) on the
repeated one — the shape an OpenAI-compatible caching upstream returns."""
long_calls = {"n": 0}
def _respond(request: Any) -> Response:
body = json.loads(request.content)
is_long = body["messages"][0]["role"] == "system"
if not is_long:
usage: dict[str, Any] = {"prompt_tokens": 5, "completion_tokens": 1}
if report_cost:
usage["cost"] = 9e-7
else:
long_calls["n"] += 1
usage = {"prompt_tokens": 3000, "completion_tokens": 1}
if long_calls["n"] > 1 and cached_tokens:
usage["prompt_tokens_details"] = {"cached_tokens": cached_tokens}
if report_cost:
usage["cost"] = 5e-5 if long_calls["n"] > 1 else 4e-4
payload = _mock_chat_response()
payload["usage"] = usage
return Response(200, json=payload)
return _respond
@pytest.mark.integration
@pytest.mark.asyncio
async def test_certify_requires_admin_auth(
@@ -205,13 +233,17 @@ async def test_certify_unknown_provider_returns_404(
async def test_certify_all_ok(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
provider_id = await _seed_and_init(integration_session, integration_client)
provider_id = await _seed_and_init(
integration_session,
integration_client,
pricing_overrides={"input_cache_read": 1.4e-8},
)
respx.get("https://certify-upstream.example/v1/models").mock(
return_value=Response(200, json=_mock_models_response())
)
respx.post("https://certify-upstream.example/v1/chat/completions").mock(
return_value=Response(200, json=_mock_chat_response())
side_effect=_caching_upstream()
)
resp = await integration_client.post(
@@ -230,6 +262,9 @@ async def test_certify_all_ok(
"endpoint.models_payload",
"usage.capture",
"cost.prompt_completion",
"cache.reported",
"cache.billing",
"cost.margin",
]
for row_id in live_row_ids:
row = _find_row(body["rows"], row_id)
@@ -239,6 +274,190 @@ async def test_certify_all_ok(
assert item["status"] == "ok", f"{item['goal']}: {item}"
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_model_path_pins_every_completion(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
base_url = "https://openrouter.ai/api/v1"
provider_id = await _seed_and_init(
integration_session,
integration_client,
pricing_overrides={"input_cache_read": 1.4e-8},
base_url=base_url,
)
model_path = encode_model_path(base_url, "cert-test-model", "azure")
integration_session.add(
ModelPathRow(
model_id="cert-test-model",
path=model_path,
provider_slug="openrouter",
provider_type="openrouter",
endpoint_tag="azure",
endpoint_name="Azure",
model_metadata="{}",
upstream_provider_id=provider_id,
)
)
await integration_session.commit()
respx.get(f"{base_url}/models").mock(
return_value=Response(200, json=_mock_models_response())
)
chat = respx.post(f"{base_url}/chat/completions").mock(
side_effect=_caching_upstream()
)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/certify",
headers=_admin_headers(),
json={"model_id": "cert-test-model", "model_path": model_path},
)
assert resp.status_code == 200, resp.text
assert chat.call_count == 3
for call in chat.calls:
body = json.loads(call.request.content)
assert body["provider"] == {
"order": ["azure"],
"allow_fallbacks": False,
}
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_uses_selected_path_pricing_for_margin(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
base_url = "https://openrouter.ai/api/v1"
provider_id = await _seed_and_init(
integration_session,
integration_client,
provider_fee=0.4,
pricing_overrides={
"prompt": 1e-7,
"completion": 5e-7,
"input_cache_read": 1e-8,
},
base_url=base_url,
)
model_path = encode_model_path(
base_url, "cert-test-model", "deepinfra/fp8"
)
integration_session.add(
ModelPathRow(
model_id="cert-test-model",
path=model_path,
provider_slug="openrouter",
provider_type="openrouter",
endpoint_tag="deepinfra/fp8",
endpoint_name="DeepInfra",
model_metadata=json.dumps(
{
"id": "cert-test-model",
"pricing": {
"prompt": 1.4e-7,
"completion": 4.2e-7,
"input_cache_read": 4.2e-9,
},
}
),
upstream_provider_id=provider_id,
)
)
await integration_session.commit()
respx.get(f"{base_url}/models").mock(
return_value=Response(200, json=_mock_models_response())
)
usages = iter(
[
{
"prompt_tokens": 31,
"completion_tokens": 1,
"cost": 0.00000459,
},
{
"prompt_tokens": 4442,
"completion_tokens": 1,
"cost": 0.00057802,
},
{
"prompt_tokens": 4442,
"completion_tokens": 1,
"prompt_tokens_details": {"cached_tokens": 4352},
"cost": 0.00005578,
},
]
)
def _respond(_request: Any) -> Response:
payload = _mock_chat_response()
payload["usage"] = next(usages)
return Response(200, json=payload)
respx.post(f"{base_url}/chat/completions").mock(side_effect=_respond)
sats_usd = 0.0008616302499999999
with patch("routstr.payment.price.sats_usd_price", return_value=sats_usd):
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/certify",
headers=_admin_headers(),
json={"model_id": "cert-test-model", "model_path": model_path},
)
assert resp.status_code == 200, resp.text
margin = _find_row(resp.json()["rows"], "cost.margin")
assert [
(sample["upstream_msats_with_fee"], sample["configured_msats"])
for sample in margin["evidence"]["samples"]
] == [(3, 3), (269, 289), (26, 15)]
assert "289 < 269" not in margin["detail"]
assert "15 < 26" in margin["detail"]
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_provider_models_includes_certification_paths(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
base_url = "https://openrouter.ai/api/v1"
provider_id = await _seed_and_init(
integration_session, integration_client, base_url=base_url
)
model_path = encode_model_path(base_url, "cert-test-model", "azure")
integration_session.add(
ModelPathRow(
model_id="cert-test-model",
path=model_path,
provider_slug="openrouter",
provider_type="openrouter",
endpoint_tag="azure",
endpoint_name="Azure",
model_metadata="{}",
upstream_provider_id=provider_id,
)
)
await integration_session.commit()
respx.get(f"{base_url}/models").mock(
return_value=Response(200, json=_mock_models_response())
)
resp = await integration_client.get(
f"/admin/api/upstream-providers/{provider_id}/models",
headers=_admin_headers(),
)
assert resp.status_code == 200, resp.text
assert resp.json()["certification_paths"]["cert-test-model"] == [
{
"path": model_path,
"endpoint_tag": "azure",
"endpoint_name": "Azure",
}
]
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
@@ -527,6 +746,76 @@ async def test_certify_with_explicit_model_id(
assert usage_row["status"] == "ok"
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_explicit_discovered_model_without_override(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
"""An operator can probe a discovered model before creating an override."""
from routstr.payment.models import (
Architecture,
Model,
Pricing,
_update_model_sats_pricing,
)
provider = await _make_provider(integration_session)
assert provider.id is not None
remote_model = _update_model_sats_pricing(
Model(
id="remote-model",
name="Remote model",
description="",
created=0,
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="unknown",
instruct_type=None,
),
pricing=Pricing(prompt=1e-7, completion=2e-7),
sats_pricing=None,
per_request_limits=None,
top_provider=None,
enabled=True,
upstream_provider_id=provider.id,
canonical_slug=None,
),
0.0005,
)
class FakeUpstream:
db_id = provider.id
def get_cached_models(self) -> list[Model]:
return [remote_model]
respx.get("https://certify-upstream.example/v1/models").mock(
return_value=Response(
200, json=_mock_models_response(models=[{"id": "remote-model"}])
)
)
respx.post("https://certify-upstream.example/v1/chat/completions").mock(
return_value=Response(200, json=_mock_chat_response(model="remote-model"))
)
with patch("routstr.proxy.get_upstreams", return_value=[FakeUpstream()]):
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider.id}/certify",
headers=_admin_headers(),
json={"model_id": "remote-model", "check_cache": False},
)
assert resp.status_code == 200, resp.text
rows = resp.json()["rows"]
assert _find_row(rows, "endpoint.reachable")["status"] == "ok"
assert _find_row(rows, "usage.capture")["status"] == "ok"
assert _find_row(rows, "cost.prompt_completion")["status"] == "ok"
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
@@ -595,3 +884,119 @@ async def test_certify_row_contract_shape(
assert set(item) >= {"goal", "label", "status", "tick", "rows"}
assert item["status"] in {"ok", "warn", "fail"}
assert item["tick"] in {"☑️", "⚠️", "❌"}
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_cache_warn_when_upstream_never_hits(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
provider_id = await _seed_and_init(integration_session, integration_client)
respx.get("https://certify-upstream.example/v1/models").mock(
return_value=Response(200, json=_mock_models_response())
)
respx.post("https://certify-upstream.example/v1/chat/completions").mock(
side_effect=_caching_upstream(cached_tokens=0)
)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/certify",
headers=_admin_headers(),
json={},
)
assert resp.status_code == 200, resp.text
rows = resp.json()["rows"]
assert _find_row(rows, "cache.reported")["status"] == "warn"
assert _find_row(rows, "cache.billing")["status"] == "warn"
goals = {item["goal"]: item["status"] for item in resp.json()["checklist"]}
assert goals["caching"] == "warn"
assert goals["margin"] == "ok"
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_cache_billing_warns_without_cache_rate(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
provider_id = await _seed_and_init(integration_session, integration_client)
respx.get("https://certify-upstream.example/v1/models").mock(
return_value=Response(200, json=_mock_models_response())
)
respx.post("https://certify-upstream.example/v1/chat/completions").mock(
side_effect=_caching_upstream(report_cost=False)
)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/certify",
headers=_admin_headers(),
json={},
)
assert resp.status_code == 200, resp.text
rows = resp.json()["rows"]
assert _find_row(rows, "cache.reported")["status"] == "ok"
billing = _find_row(rows, "cache.billing")
assert billing["status"] == "warn"
assert billing["evidence"]["actual_total_msats"] == 841
assert _find_row(rows, "cost.margin")["status"] == "warn"
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_margin_fails_when_upstream_costs_more(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
provider_id = await _seed_and_init(integration_session, integration_client)
respx.get("https://certify-upstream.example/v1/models").mock(
return_value=Response(200, json=_mock_models_response())
)
def _expensive(request: Any) -> Response:
payload = _mock_chat_response()
payload["usage"] = {"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3}
return Response(200, json=payload)
respx.post("https://certify-upstream.example/v1/chat/completions").mock(
side_effect=_expensive
)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/certify",
headers=_admin_headers(),
json={},
)
assert resp.status_code == 200, resp.text
margin = _find_row(resp.json()["rows"], "cost.margin")
assert margin["status"] == "fail"
goals = {item["goal"]: item["status"] for item in resp.json()["checklist"]}
assert goals["margin"] == "fail"
@pytest.mark.integration
@pytest.mark.asyncio
@respx.mock
async def test_certify_check_cache_false_skips_probe(
integration_client: AsyncClient, integration_session: AsyncSession
) -> None:
provider_id = await _seed_and_init(integration_session, integration_client)
respx.get("https://certify-upstream.example/v1/models").mock(
return_value=Response(200, json=_mock_models_response())
)
chat = respx.post("https://certify-upstream.example/v1/chat/completions").mock(
return_value=Response(200, json=_mock_chat_response())
)
resp = await integration_client.post(
f"/admin/api/upstream-providers/{provider_id}/certify",
headers=_admin_headers(),
json={"check_cache": False},
)
assert resp.status_code == 200, resp.text
assert chat.call_count == 1
rows = resp.json()["rows"]
for row_id in ("cache.reported", "cache.billing", "cost.margin"):
row = _find_row(rows, row_id)
assert row["status"] == "warn"
assert "disabled" in row["detail"]
+4 -1
View File
@@ -527,9 +527,12 @@ class TestBuildChecklist:
self._row("cost.prompt_completion", STATUS_OK),
self._row("pricing.served_matches_configured", STATUS_OK),
self._row("pricing.enabled_models_served", STATUS_OK),
self._row("cache.reported", STATUS_OK),
self._row("cache.billing", STATUS_OK),
self._row("cost.margin", STATUS_OK),
]
checklist = build_checklist(rows)
assert len(checklist) == 4
assert len(checklist) == 6
for item in checklist:
assert item["status"] == STATUS_OK
assert item["tick"] == "☑️"
+499
View File
@@ -0,0 +1,499 @@
"""Unit tests for the cache and margin rows in
routstr.upstream.certification_cache — no DB, network only via respx."""
from __future__ import annotations
import json
from typing import Any
import httpx
import pytest
import respx
from routstr.payment.cost_calculation import CostData, CostDataError
from routstr.upstream.certification import STATUS_FAIL, STATUS_OK, STATUS_WARN
from routstr.upstream.certification_cache import (
CacheProbeResult,
_raw_cache_keys,
cache_billing_row,
cache_probe_prefix,
cache_reported_row,
cost_margin_row,
probe_cache,
skipped_cache_rows,
)
SATS_USD = 0.0005
CHAT_URL = "https://upstream.example/v1/chat/completions"
def _model(
prompt: float = 1.4e-7,
completion: float = 2.8e-7,
cache_read: float = 0.0,
sats_usd: float = SATS_USD,
) -> Any:
from routstr.payment.models import (
Architecture,
Model,
Pricing,
_update_model_sats_pricing,
)
model = Model(
id="test-model",
name="test-model",
created=0,
description="",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="unknown",
instruct_type=None,
),
pricing=Pricing(
prompt=prompt, completion=completion, input_cache_read=cache_read
),
sats_pricing=None,
per_request_limits=None,
top_provider=None,
enabled=True,
upstream_provider_id=None,
canonical_slug=None,
)
return _update_model_sats_pricing(model, sats_usd)
def _payload(usage: dict[str, Any] | None) -> dict[str, Any]:
body: dict[str, Any] = {"choices": [{"message": {"content": "ok"}}]}
if usage is not None:
body["usage"] = usage
return body
def _probe(
payloads: list[dict[str, Any] | None],
statuses: list[int | None] | None = None,
) -> CacheProbeResult:
statuses = statuses if statuses is not None else [200] * len(payloads)
return CacheProbeResult(
chat_url=CHAT_URL,
statuses=statuses,
payloads=payloads,
errors=[None] * len(payloads),
latencies_ms=[1.0] * len(payloads),
)
def _cost(total: int) -> CostData:
return CostData(base_msats=0, input_msats=total, output_msats=0, total_msats=total)
CACHED = {
"prompt_tokens": 3000,
"completion_tokens": 1,
"prompt_tokens_details": {"cached_tokens": 2900},
}
UNCACHED = {"prompt_tokens": 3000, "completion_tokens": 1}
class TestPrefix:
def test_prefix_is_long_and_deterministic(self) -> None:
prefix = cache_probe_prefix()
assert len(prefix) > 8000
assert prefix == cache_probe_prefix()
class TestRawCacheKeys:
def test_finds_nested_positive_cache_fields(self) -> None:
usage = {"prompt_tokens": 5, "details": {"cache_hits": 3, "cached": 0}}
assert _raw_cache_keys(usage) == ["details.cache_hits"]
def test_ignores_bool_and_non_numeric(self) -> None:
assert _raw_cache_keys({"cached": True, "cache_key": "abc"}) == []
class TestCacheReportedRow:
def test_ok_when_second_call_reports_cached_tokens(self) -> None:
row = cache_reported_row(_probe([_payload(UNCACHED), _payload(CACHED)]))
assert row["status"] == STATUS_OK
assert row["evidence"]["second_usage"]["cache_read_tokens"] == 2900
def test_warn_when_no_cache_hit(self) -> None:
row = cache_reported_row(_probe([_payload(UNCACHED), _payload(UNCACHED)]))
assert row["status"] == STATUS_WARN
def test_fail_when_cache_reported_under_unknown_field(self) -> None:
second = _payload(
{"prompt_tokens": 3000, "completion_tokens": 1, "cache_hit_tokens": 2900}
)
row = cache_reported_row(_probe([_payload(UNCACHED), second]))
assert row["status"] == STATUS_FAIL
assert row["evidence"]["unrecognised_cache_fields"] == ["cache_hit_tokens"]
def test_fail_when_first_call_failed(self) -> None:
probe = CacheProbeResult(
chat_url=CHAT_URL,
statuses=[None],
payloads=[None],
errors=["ConnectError: boom"],
latencies_ms=[1.0],
)
row = cache_reported_row(probe)
assert row["status"] == STATUS_FAIL
assert "ConnectError" in row["detail"]
def test_fail_when_second_call_non_2xx(self) -> None:
row = cache_reported_row(
_probe([_payload(UNCACHED), {"error": "rate limited"}], [200, 429])
)
assert row["status"] == STATUS_FAIL
class TestCacheBillingRow:
def test_warn_when_nothing_cached(self) -> None:
row = cache_billing_row(
model=_model(), probe=_probe([_payload(UNCACHED)] * 2), cost_data=_cost(1)
)
assert row["status"] == STATUS_WARN
def test_warn_when_pricing_unknown(self) -> None:
row = cache_billing_row(
model=_model(),
probe=_probe([_payload(UNCACHED), _payload(CACHED)]),
cost_data=_cost(1),
pricing_known=False,
)
assert row["status"] == STATUS_WARN
def test_fail_on_engine_error(self) -> None:
row = cache_billing_row(
model=_model(),
probe=_probe([_payload(UNCACHED), _payload(CACHED)]),
cost_data=CostDataError(message="nope", code="pricing_error"),
)
assert row["status"] == STATUS_FAIL
def test_ok_with_discounted_rate(self) -> None:
# 280 msats/1k input, 28 msats/1k cached, 560 msats/1k output:
# 100*0.28 + 2900*0.028 + 1*0.56 = 109.76 -> 110
row = cache_billing_row(
model=_model(cache_read=1.4e-8),
probe=_probe([_payload(UNCACHED), _payload(CACHED)]),
cost_data=_cost(110),
)
assert row["status"] == STATUS_OK, row
assert row["evidence"]["full_price_total_msats"] == 841
def test_warn_when_cached_billed_at_full_rate(self) -> None:
row = cache_billing_row(
model=_model(),
probe=_probe([_payload(UNCACHED), _payload(CACHED)]),
cost_data=_cost(841),
)
assert row["status"] == STATUS_WARN
assert "full input rate" in row["detail"]
def test_fail_when_engine_disagrees(self) -> None:
row = cache_billing_row(
model=_model(cache_read=1.4e-8),
probe=_probe([_payload(UNCACHED), _payload(CACHED)]),
cost_data=_cost(841),
)
assert row["status"] == STATUS_FAIL
def test_ok_when_upstream_reports_cost(self) -> None:
cached = _payload({**CACHED, "cost": 5e-5})
row = cache_billing_row(
model=_model(),
probe=_probe([_payload(UNCACHED), cached]),
cost_data=_cost(100),
)
assert row["status"] == STATUS_OK
assert row["evidence"]["reported_usd"] == 5e-5
class TestCostMarginRow:
def _row(
self, payloads: list[Any], fee: float = 1.0, model: Any = None
) -> dict[str, Any]:
return cost_margin_row(
model=model or _model(),
payloads=payloads,
provider_fee=fee,
sats_to_usd=SATS_USD,
)
def test_warn_when_no_cost_reported(self) -> None:
row = self._row([_payload(UNCACHED), None])
assert row["status"] == STATUS_WARN
assert row["evidence"]["samples"] == []
def test_ok_when_configured_covers_upstream(self) -> None:
# configured: 5*0.28 + 1*0.56 = 1.96 -> 2 msats; upstream 9e-7 USD -> 2
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
row = self._row([payload])
assert row["status"] == STATUS_OK, row
assert row["evidence"]["samples"][0]["upstream_msats_with_fee"] == 2
def test_fail_when_upstream_costs_more(self) -> None:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3})
row = self._row([payload])
assert row["status"] == STATUS_FAIL
assert "lose money" in row["detail"]
def test_fee_scales_upstream_cost(self) -> None:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
assert self._row([payload], fee=1.0)["status"] == STATUS_OK
assert self._row([payload], fee=2.0)["status"] == STATUS_FAIL
def test_deepseek_endpoint_price_exceeds_configured_model_price(self) -> None:
sats_usd = 0.0008616302499999999
model = _model(
prompt=4e-8,
completion=2e-7,
cache_read=4e-9,
sats_usd=sats_usd,
)
payloads = [
_payload(
{
"prompt_tokens": 31,
"completion_tokens": 1,
"cost": 0.00000459,
}
),
_payload(
{
"prompt_tokens": 4442,
"completion_tokens": 1,
"cost": 0.00057802,
}
),
_payload(
{
"prompt_tokens": 4442,
"completion_tokens": 1,
"prompt_tokens_details": {"cached_tokens": 4352},
"cost": 0.00005578,
}
),
]
row = cost_margin_row(
model=model,
payloads=payloads,
provider_fee=0.4,
sats_to_usd=sats_usd,
)
assert row["status"] == STATUS_FAIL
assert [
(sample["upstream_msats_with_fee"], sample["configured_msats"])
for sample in row["evidence"]["samples"]
] == [(3, 2), (269, 207), (26, 25)]
assert "207 < 269" in row["detail"]
assert "2 < 3" not in row["detail"]
assert "25 < 26" not in row["detail"]
def test_one_msat_margin_gap_is_rounding_tolerance(self) -> None:
sats_usd = 0.0008616302499999999
model = _model(
prompt=4e-8,
completion=2e-7,
cache_read=4e-9,
sats_usd=sats_usd,
)
payloads = [
_payload(
{
"prompt_tokens": 31,
"completion_tokens": 1,
"cost": 0.00000459,
}
),
_payload(
{
"prompt_tokens": 4442,
"completion_tokens": 1,
"prompt_tokens_details": {"cached_tokens": 4352},
"cost": 0.00005578,
}
),
]
row = cost_margin_row(
model=model,
payloads=payloads,
provider_fee=0.4,
sats_to_usd=sats_usd,
)
assert row["status"] == STATUS_OK
def test_warn_when_pricing_unknown(self) -> None:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
row = cost_margin_row(
model=_model(),
payloads=[payload],
provider_fee=1.0,
sats_to_usd=SATS_USD,
pricing_known=False,
)
assert row["status"] == STATUS_WARN
class TestSkippedRows:
def test_three_warn_rows(self) -> None:
rows = skipped_cache_rows("Skipped")
assert [r["id"] for r in rows] == [
"cache.reported",
"cache.billing",
"cost.margin",
]
assert all(r["status"] == STATUS_WARN for r in rows)
class TestProbeCache:
@pytest.mark.asyncio
@respx.mock
async def test_sends_same_prompt_twice_with_cache_control(self) -> None:
route = respx.post(CHAT_URL).mock(
return_value=httpx.Response(200, json=_payload(CACHED))
)
async with httpx.AsyncClient() as client:
result = await probe_cache(
"https://upstream.example/v1", "k", "m", client=client
)
assert result.request_format == "cache_control"
assert route.call_count == 2
first, second = (json.loads(c.request.content) for c in route.calls)
assert first == second
system = first["messages"][0]["content"]
assert system[0]["cache_control"] == {"type": "ephemeral"}
assert first["max_tokens"] == 1
assert route.calls[0].request.headers["Authorization"] == "Bearer k"
@pytest.mark.asyncio
@respx.mock
async def test_pins_both_requests_to_one_endpoint(self) -> None:
route = respx.post(CHAT_URL).mock(
return_value=httpx.Response(200, json=_payload(CACHED))
)
async with httpx.AsyncClient() as client:
await probe_cache(
"https://upstream.example/v1",
"",
"m",
endpoint_tag="azure/swedencentral",
client=client,
)
assert route.call_count == 2
for call in route.calls:
body = json.loads(call.request.content)
assert body["provider"] == {
"order": ["azure/swedencentral"],
"allow_fallbacks": False,
}
@pytest.mark.asyncio
@respx.mock
async def test_falls_back_to_plain_system_on_rejection(self) -> None:
route = respx.post(CHAT_URL).mock(
side_effect=[
httpx.Response(400, json={"error": "cache_control not allowed"}),
httpx.Response(200, json=_payload(UNCACHED)),
httpx.Response(200, json=_payload(CACHED)),
]
)
async with httpx.AsyncClient() as client:
result = await probe_cache(
"https://upstream.example/v1", "", "m", client=client
)
assert result.request_format == "plain"
assert route.call_count == 3
assert result.statuses == [200, 200]
body = json.loads(route.calls[1].request.content)
assert isinstance(body["messages"][0]["content"], str)
@pytest.mark.asyncio
@respx.mock
async def test_stops_after_failed_first_call(self) -> None:
route = respx.post(CHAT_URL).mock(side_effect=httpx.ConnectError("boom"))
async with httpx.AsyncClient() as client:
result = await probe_cache(
"https://upstream.example/v1", "", "m", client=client
)
assert route.call_count == 1
assert result.payloads == [None]
assert "ConnectError" in (result.second_error or "")
class TestErrorBranches:
@pytest.mark.asyncio
@respx.mock
async def test_non_json_and_list_bodies_are_errors(self) -> None:
respx.post(CHAT_URL).mock(
side_effect=[
httpx.Response(200, content=b"not json"),
httpx.Response(200, json=[1, 2]),
]
)
async with httpx.AsyncClient() as client:
result = await probe_cache(
"https://upstream.example/v1", "", "m", client=client
)
assert result.payloads == [None, None]
assert result.errors[0] is not None
assert "list" in (result.errors[1] or "")
def test_malformed_usage_object_is_not_a_hit(self) -> None:
second = _payload({"prompt_tokens": {"nested": True}})
row = cache_reported_row(_probe([_payload(UNCACHED), second]))
assert row["status"] == STATUS_WARN
def test_billing_fails_on_non_finite_rate(self) -> None:
row = cache_billing_row(
model=_model(prompt=float("inf"), cache_read=1.4e-8),
probe=_probe([_payload(UNCACHED), _payload(CACHED)]),
cost_data=_cost(1),
)
assert row["status"] == STATUS_FAIL
assert "error" in row["evidence"]
def test_margin_fails_on_zero_sats_price(self) -> None:
payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7})
row = cost_margin_row(
model=_model(), payloads=[payload], provider_fee=1.0, sats_to_usd=0.0
)
assert row["status"] == STATUS_FAIL
@pytest.mark.asyncio
async def test_engine_raise_becomes_billing_fail(self) -> None:
from unittest.mock import patch
from routstr.upstream.certification_cache import run_cache_checks
async def _boom(*args: Any, **kwargs: Any) -> Any:
raise RuntimeError("engine exploded")
with (
patch(
"routstr.upstream.certification_cache.probe_cache",
return_value=_probe([_payload(UNCACHED), _payload(CACHED)]),
),
patch("routstr.upstream.certification_cache.calculate_cost", _boom),
):
rows = await run_cache_checks(
"https://upstream.example/v1",
"",
_model(),
provider_fee=1.0,
sats_to_usd=SATS_USD,
probe_payload=_payload(UNCACHED),
)
by_id = {r["id"]: r for r in rows}
assert by_id["cache.billing"]["status"] == STATUS_FAIL
assert "engine exploded" in by_id["cache.billing"]["detail"]
+2 -2
View File
@@ -337,8 +337,8 @@ class TestCliFreshProcess:
)
document = json.loads(out.read_text(encoding="utf-8"))
assert isinstance(document, list) and document
assert len(document[0]["rows"]) == 5
assert len(document[0]["checklist"]) == 4
assert len(document[0]["rows"]) == 8
assert len(document[0]["checklist"]) == 6
def test_dead_host_exits_non_zero(self) -> None:
result = _run_cli("--url", "http://localhost:1/v1", "--timeout", "1")
+20
View File
@@ -25,9 +25,11 @@ import {
AlertTriangle,
Unlock,
Loader2,
ShieldCheck,
} from 'lucide-react';
import { ProviderBalance } from '@/components/provider-balance';
import { ProviderModelsPanel } from '@/components/provider-models-panel';
import { ProviderCertificationDialog } from '@/components/provider-certification-dialog';
import { RoutstrCreateKeySection } from '@/components/providers/RoutstrCreateKeySection';
import { RoutstrProviderService } from '@/lib/api/services/routstr-provider';
import { getErrorStatus } from '@/lib/api/client';
@@ -93,6 +95,7 @@ export function ProviderCard({
}: ProviderCardProps) {
const queryClient = useQueryClient();
const [isKeyModalOpen, setIsKeyModalOpen] = useState(false);
const [isCertifyOpen, setIsCertifyOpen] = useState(false);
const [isReleaseDialogOpen, setIsReleaseDialogOpen] = useState(false);
// The claim as the query cache held it when the admin opened the dialog.
// The mutation sends this token rather than re-reading the query at submit
@@ -310,6 +313,17 @@ export function ProviderCard({
)}
</Button>
<Button
variant='outline'
size='sm'
onClick={() => setIsCertifyOpen(true)}
className='justify-center gap-1.5'
title='Probe the upstream and verify usage, pricing, caching and margin'
>
<ShieldCheck className='h-4 w-4' />
<span>Certify</span>
</Button>
<Button
variant='outline'
size='sm'
@@ -369,6 +383,12 @@ export function ProviderCard({
</AlertDialogContent>
</AlertDialog>
<ProviderCertificationDialog
provider={provider}
open={isCertifyOpen}
onOpenChange={setIsCertifyOpen}
/>
<Dialog open={isKeyModalOpen} onOpenChange={setIsKeyModalOpen}>
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[500px]'>
<DialogHeader>
@@ -0,0 +1,697 @@
'use client';
import { useEffect, useState } from 'react';
import { useMutation, useQuery } from '@tanstack/react-query';
import {
AlertTriangle,
CheckCircle2,
ChevronDown,
ChevronRight,
Loader2,
RotateCcw,
XCircle,
} from 'lucide-react';
import { AdminService } from '@/lib/api/services/admin';
import type {
AdminModel,
CertificationPath,
CertificationRow,
CertificationStatus,
ProviderCertification,
UpstreamProvider,
} from '@/lib/api/services/admin';
import { Badge } from '@/components/ui/badge';
import { Button } from '@/components/ui/button';
import { Checkbox } from '@/components/ui/checkbox';
import {
Command,
CommandEmpty,
CommandGroup,
CommandInput,
CommandItem,
CommandList,
} from '@/components/ui/command';
import {
Dialog,
DialogContent,
DialogDescription,
DialogHeader,
DialogTitle,
} from '@/components/ui/dialog';
import { Label } from '@/components/ui/label';
import {
Popover,
PopoverContent,
PopoverTrigger,
} from '@/components/ui/popover';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import { ToggleGroup, ToggleGroupItem } from '@/components/ui/toggle-group';
import { cn } from '@/lib/utils';
interface ProviderCertificationDialogProps {
provider: UpstreamProvider;
open: boolean;
onOpenChange: (open: boolean) => void;
}
interface ModelOption {
model: AdminModel;
source: 'configured' | 'discovered';
}
type ModelPathMode = 'default' | 'selected' | 'all';
interface ModelPathTarget {
path?: string;
label: string;
}
interface ModelRun {
modelId: string;
targets: ModelPathTarget[];
}
interface ModelCertificationResult {
resultKey: string;
modelId: string;
pathLabel: string;
report?: ProviderCertification;
error?: string;
}
const STATUS_STYLES: Record<
CertificationStatus,
{ label: string; icon: typeof CheckCircle2; className: string }
> = {
ok: {
label: 'OK',
icon: CheckCircle2,
className:
'border-emerald-500/40 bg-emerald-500/10 text-emerald-700 dark:text-emerald-400',
},
warn: {
label: 'Warn',
icon: AlertTriangle,
className:
'border-amber-500/40 bg-amber-500/10 text-amber-700 dark:text-amber-400',
},
fail: {
label: 'Fail',
icon: XCircle,
className: 'border-red-500/40 bg-red-500/10 text-red-700 dark:text-red-400',
},
};
function StatusBadge({ status }: { status: CertificationStatus }) {
const style = STATUS_STYLES[status];
const Icon = style.icon;
return (
<Badge variant='outline' className={cn('gap-1', style.className)}>
<Icon className='h-3 w-3' />
{style.label}
</Badge>
);
}
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 ChecklistSummary({ report }: { report: ProviderCertification }) {
return (
<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>
);
}
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'>
<ChecklistSummary report={report} />
<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>
);
}
function getErrorMessage(error: unknown): string {
if (error instanceof Error) return error.message;
return 'Certification request failed';
}
function resultStatus(
result: ModelCertificationResult
): CertificationStatus | 'error' {
if (result.error || !result.report) return 'error';
if (result.report.rows.some((row) => row.status === 'fail')) return 'fail';
if (result.report.rows.some((row) => row.status === 'warn')) return 'warn';
return 'ok';
}
export function ProviderCertificationDialog({
provider,
open,
onOpenChange,
}: ProviderCertificationDialogProps) {
const [checkCache, setCheckCache] = useState(true);
const [selectedModelIds, setSelectedModelIds] = useState<string[]>([]);
const [pathModes, setPathModes] = useState<Record<string, ModelPathMode>>({});
const [selectedModelPaths, setSelectedModelPaths] = useState<
Record<string, string[]>
>({});
const [results, setResults] = useState<ModelCertificationResult[]>([]);
const [currentModel, setCurrentModel] = useState<{
id: string;
index: number;
total: number;
pathCount: number;
} | null>(null);
const models = useQuery({
queryKey: ['provider-models', provider.id],
queryFn: () => AdminService.getProviderModels(provider.id),
enabled: open,
});
const certify = useMutation({
mutationFn: async ({
modelRuns,
includeCache,
}: {
modelRuns: ModelRun[];
includeCache: boolean;
}) => {
const completed: ModelCertificationResult[] = [];
setResults([]);
for (const [index, run] of modelRuns.entries()) {
setCurrentModel({
id: run.modelId,
index: index + 1,
total: modelRuns.length,
pathCount: run.targets.length,
});
const batch = await Promise.all(
run.targets.map(async (target): Promise<ModelCertificationResult> => {
const resultKey = `${run.modelId}::${target.path ?? 'default'}`;
try {
const report = await AdminService.certifyProvider(provider.id, {
model_id: run.modelId,
model_path: target.path,
check_cache: includeCache,
});
return {
resultKey,
modelId: run.modelId,
pathLabel: target.label,
report,
};
} catch (error) {
return {
resultKey,
modelId: run.modelId,
pathLabel: target.label,
error: getErrorMessage(error),
};
}
})
);
completed.push(...batch);
setResults([...completed]);
}
return completed;
},
onSettled: () => setCurrentModel(null),
});
const { reset: resetCertification } = certify;
useEffect(() => {
if (!open) {
resetCertification();
setSelectedModelIds([]);
setPathModes({});
setSelectedModelPaths({});
setResults([]);
setCurrentModel(null);
}
}, [open, resetCertification]);
const configuredOptions: ModelOption[] =
models.data?.db_models.map((model) => ({
model,
source: 'configured',
})) ?? [];
const discoveredOptions: ModelOption[] =
models.data?.remote_models.map((model) => ({
model,
source: 'discovered',
})) ?? [];
const allOptions = [...configuredOptions, ...discoveredOptions];
const namesById = new Map(
allOptions.map(({ model }) => [model.id, model.name || model.id])
);
const pathsForModel = (modelId: string): CertificationPath[] =>
(models.data?.certification_paths[modelId] ?? []).filter(
(path) => path.endpoint_tag
);
const pathLabel = (path: CertificationPath): string =>
path.endpoint_name && path.endpoint_name !== path.endpoint_tag
? `${path.endpoint_name} (${path.endpoint_tag})`
: path.endpoint_tag || 'Provider default';
const toggleModel = (modelId: string) => {
const isSelected = selectedModelIds.includes(modelId);
setSelectedModelIds((current) =>
isSelected
? current.filter((id) => id !== modelId)
: [...current, modelId]
);
setPathModes((current) => {
const next = { ...current };
if (isSelected) delete next[modelId];
else next[modelId] = 'default';
return next;
});
setSelectedModelPaths((current) => {
const next = { ...current };
if (isSelected) delete next[modelId];
return next;
});
resetCertification();
setResults([]);
};
const modelsNeedingPath = selectedModelIds.filter(
(modelId) =>
pathModes[modelId] === 'selected' &&
(selectedModelPaths[modelId]?.length ?? 0) === 0
);
const buildModelRuns = (): ModelRun[] =>
selectedModelIds.map((modelId) => {
const paths = pathsForModel(modelId);
const mode = pathModes[modelId] ?? 'default';
if (mode === 'all') {
return {
modelId,
targets: paths.map((path) => ({
path: path.path,
label: pathLabel(path),
})),
};
}
if (mode === 'selected') {
const selected = new Set(selectedModelPaths[modelId] ?? []);
return {
modelId,
targets: paths
.filter((path) => selected.has(path.path))
.map((path) => ({ path: path.path, label: pathLabel(path) })),
};
}
return { modelId, targets: [{ label: 'Provider default' }] };
});
const modelRuns = buildModelRuns();
const targetCount = modelRuns.reduce(
(total, run) => total + run.targets.length,
0
);
const renderModelGroup = (label: string, options: ModelOption[]) => {
if (options.length === 0) return null;
return (
<CommandGroup heading={`${label} (${options.length})`}>
{options.map(({ model }) => (
<CommandItem
key={`${label}-${model.id}`}
value={`${model.name} ${model.id}`}
onSelect={() => toggleModel(model.id)}
disabled={certify.isPending}
>
<Checkbox
checked={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 (
<Dialog open={open} onOpenChange={onOpenChange}>
<DialogContent className='max-h-[90dvh] overflow-y-auto sm:max-w-[780px]'>
<DialogHeader>
<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>
<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'>
{selectedModelIds.length} selected
</span>
{selectedModelIds.length > 0 && !certify.isPending && (
<Button
type='button'
variant='ghost'
size='sm'
onClick={() => {
setSelectedModelIds([]);
setPathModes({});
setSelectedModelPaths({});
setResults([]);
resetCertification();
}}
>
Clear
</Button>
)}
</div>
</div>
<Command className='h-auto rounded-md border'>
<CommandInput
placeholder='Search models by name or ID…'
disabled={models.isLoading || certify.isPending}
/>
<CommandList className='max-h-[min(18rem,40dvh)]'>
<CommandEmpty>
{models.isLoading ? 'Loading models…' : 'No models found'}
</CommandEmpty>
{renderModelGroup('Configured models', configuredOptions)}
{renderModelGroup('Discovered models', discoveredOptions)}
</CommandList>
</Command>
{models.isError && (
<p className='text-destructive text-sm'>
{getErrorMessage(models.error)}
</p>
)}
{selectedModelIds.map((modelId) => {
const paths = pathsForModel(modelId);
const mode = pathModes[modelId] ?? 'default';
const selectedPaths = 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;
setPathModes((current) => ({
...current,
[modelId]: value as ModelPathMode,
}));
}}
disabled={certify.isPending}
className='w-full justify-start'
>
<ToggleGroupItem value='default'>Default</ToggleGroupItem>
<ToggleGroupItem value='selected'>Choose paths</ToggleGroupItem>
<ToggleGroupItem value='all'>All paths</ToggleGroupItem>
</ToggleGroup>
{mode === 'default' && (
<p className='text-muted-foreground text-xs'>
Uses the upstream provider&apos;s normal model routing.
</p>
)}
{mode === 'selected' && (
<Popover>
<PopoverTrigger asChild>
<Button
type='button'
variant='outline'
size='sm'
className='w-full justify-between'
disabled={certify.isPending}
>
<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) =>
setSelectedModelPaths((current) => {
const previous = current[modelId] ?? [];
return {
...current,
[modelId]:
value === true
? [...previous, path.path]
: previous.filter(
(item) => item !== path.path
),
};
})
}
/>
<span className='min-w-0 truncate'>
{pathLabel(path)}
</span>
</label>
);
})}
</div>
</PopoverContent>
</Popover>
)}
{mode === 'all' && (
<p className='text-muted-foreground text-xs'>
All {paths.length} paths will run in parallel.
</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>
<div className='flex flex-wrap items-center justify-between gap-3'>
<div className='flex items-center gap-2'>
<Checkbox
id={`certify-cache-${provider.id}`}
checked={checkCache}
onCheckedChange={(value) => setCheckCache(value === true)}
disabled={certify.isPending}
/>
<Label htmlFor={`certify-cache-${provider.id}`} className='text-sm'>
Probe prompt caching and margin
</Label>
</div>
<Button
variant='outline'
size='sm'
onClick={() =>
certify.mutate({
modelRuns,
includeCache: checkCache,
})
}
disabled={
certify.isPending ||
selectedModelIds.length === 0 ||
modelsNeedingPath.length > 0
}
className='gap-1.5'
>
{certify.isPending ? (
<Loader2 className='h-4 w-4 animate-spin' />
) : (
<RotateCcw className='h-4 w-4' />
)}
{certify.isPending
? `Running ${currentModel?.index ?? 1} of ${currentModel?.total ?? selectedModelIds.length}`
: results.length > 0
? `Run ${targetCount} route${targetCount === 1 ? '' : 's'} again`
: `Certify ${targetCount || ''} route${targetCount === 1 ? '' : 's'}`}
</Button>
</div>
{currentModel && (
<div className='text-muted-foreground flex items-center gap-2 text-sm'>
<Loader2 className='h-4 w-4 animate-spin' />
Probing {namesById.get(currentModel.id) ?? currentModel.id}
{currentModel.pathCount > 1
? ` across ${currentModel.pathCount} paths in parallel`
: ''}{' '}
— model {currentModel.index} of {currentModel.total}
</div>
)}
{results.length > 0 && (
<Tabs
key={results.map((result) => result.resultKey).join('|')}
defaultValue={results[0].resultKey}
className='space-y-3'
>
<div className='overflow-x-auto'>
<TabsList variant='line' className='min-w-max'>
{results.map((result) => {
const status = resultStatus(result);
const Icon =
status === 'error' ? XCircle : STATUS_STYLES[status].icon;
return (
<TabsTrigger
key={result.resultKey}
value={result.resultKey}
title={`${namesById.get(result.modelId) ?? result.modelId} · ${result.pathLabel}`}
className='max-w-64'
>
<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} ·{' '}
{result.pathLabel}
</span>
</TabsTrigger>
);
})}
</TabsList>
</div>
{results.map((result) => (
<TabsContent key={result.resultKey} value={result.resultKey}>
{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>
))}
</Tabs>
)}
</DialogContent>
</Dialog>
);
}
+59
View File
@@ -46,6 +46,43 @@ export const UpdateUpstreamProviderSchema = z.object({
slug: z.string().optional(),
});
export const CertificationStatusSchema = z.enum(['ok', 'warn', 'fail']);
export const CertificationRowSchema = z.object({
id: z.string(),
status: CertificationStatusSchema,
title: z.string(),
detail: z.string(),
evidence: z.record(z.string(), z.unknown()),
});
export const CertificationGoalSchema = z.object({
goal: z.string(),
label: z.string(),
status: CertificationStatusSchema,
tick: z.string(),
rows: z.array(z.string()),
});
export const ProviderCertificationSchema = z.object({
provider_id: z.number(),
generated_at: z.string(),
rows: z.array(CertificationRowSchema),
checklist: z.array(CertificationGoalSchema),
});
export type CertificationStatus = z.infer<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;
timeout_seconds?: number;
check_cache?: boolean;
};
export const AdminModelPricingSchema = z.object({
prompt: z.number().optional(),
completion: z.number().optional(),
@@ -84,6 +121,12 @@ export const AdminModelSchema = z.object({
forwarded_model_id: z.string().nullable().optional(),
});
export const CertificationPathSchema = z.object({
path: z.string(),
endpoint_tag: z.string().nullable(),
endpoint_name: z.string().nullable(),
});
export const ProviderModelsSchema = z.object({
provider: z.object({
id: z.number(),
@@ -92,6 +135,10 @@ export const ProviderModelsSchema = z.object({
}),
db_models: z.array(AdminModelSchema),
remote_models: z.array(AdminModelSchema),
certification_paths: z.record(
z.string(),
z.array(CertificationPathSchema)
),
});
export type ProviderType = z.infer<typeof ProviderTypeSchema>;
@@ -107,6 +154,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 +365,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`