mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
feat: add model path certification
This commit is contained in:
+93
-1
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, {}),
|
||||
]
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"] == "☑️"
|
||||
|
||||
@@ -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"]
|
||||
@@ -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")
|
||||
|
||||
@@ -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'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>
|
||||
);
|
||||
}
|
||||
@@ -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`
|
||||
|
||||
Reference in New Issue
Block a user