From 721eb4ff8d050ef00270422fb169177454c915d7 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 21:19:36 +0200 Subject: [PATCH] feat: add model path certification --- routstr/core/admin.py | 94 ++- routstr/upstream/certification.py | 70 +- routstr/upstream/certification_cache.py | 586 +++++++++++++++ routstr/upstream/model_paths.py | 52 ++ tests/integration/test_certify_endpoint.py | 411 ++++++++++- tests/unit/test_certification.py | 5 +- tests/unit/test_certification_cache.py | 499 +++++++++++++ tests/unit/test_certification_hardening.py | 4 +- ui/components/provider-card.tsx | 20 + .../provider-certification-dialog.tsx | 697 ++++++++++++++++++ ui/lib/api/services/admin.ts | 59 ++ 11 files changed, 2489 insertions(+), 8 deletions(-) create mode 100644 routstr/upstream/certification_cache.py create mode 100644 tests/unit/test_certification_cache.py create mode 100644 ui/components/provider-certification-dialog.tsx diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 8a51d9c6..bd5b4e18 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -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 diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index a032eedb..5a87c8d0 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -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 diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py new file mode 100644 index 00000000..3ee74e7a --- /dev/null +++ b/routstr/upstream/certification_cache.py @@ -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, {}), + ] diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index fedb042e..0ac4e09d 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -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: diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index 69f432aa..c7aa9151 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -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"] diff --git a/tests/unit/test_certification.py b/tests/unit/test_certification.py index dc3ae9fc..1e963a79 100644 --- a/tests/unit/test_certification.py +++ b/tests/unit/test_certification.py @@ -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"] == "☑️" diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py new file mode 100644 index 00000000..2ca56482 --- /dev/null +++ b/tests/unit/test_certification_cache.py @@ -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"] diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py index 52f2fb02..7eb05552 100644 --- a/tests/unit/test_certification_hardening.py +++ b/tests/unit/test_certification_hardening.py @@ -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") diff --git a/ui/components/provider-card.tsx b/ui/components/provider-card.tsx index 44f05183..ef07013a 100644 --- a/ui/components/provider-card.tsx +++ b/ui/components/provider-card.tsx @@ -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({ )} + + + {showEvidence && ( +
+              {JSON.stringify(row.evidence, null, 2)}
+            
+ )} + + )} + + ); +} + +function ChecklistSummary({ report }: { report: ProviderCertification }) { + return ( + + ); +} + +function CertificationReport({ report }: { report: ProviderCertification }) { + const failing = report.rows.filter((row) => row.status === 'fail').length; + const warning = report.rows.filter((row) => row.status === 'warn').length; + + return ( +
+ +
+ + {report.rows.length} checks · {failing} failed · {warning} warnings + + Generated {new Date(report.generated_at).toLocaleString()} +
+
    + {report.rows.map((row) => ( + + ))} +
+
+ ); +} + +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([]); + const [pathModes, setPathModes] = useState>({}); + const [selectedModelPaths, setSelectedModelPaths] = useState< + Record + >({}); + const [results, setResults] = useState([]); + 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 => { + 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 ( + + {options.map(({ model }) => ( + toggleModel(model.id)} + disabled={certify.isPending} + > + + ))} + + ); + }; + + return ( + + + + Certify upstream models + + 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}. + + + +
+
+ +
+ + {selectedModelIds.length} selected + + {selectedModelIds.length > 0 && !certify.isPending && ( + + )} +
+
+ + + + + {models.isLoading ? 'Loading models…' : 'No models found'} + + {renderModelGroup('Configured models', configuredOptions)} + {renderModelGroup('Discovered models', discoveredOptions)} + + + {models.isError && ( +

+ {getErrorMessage(models.error)} +

+ )} + {selectedModelIds.map((modelId) => { + const paths = pathsForModel(modelId); + const mode = pathModes[modelId] ?? 'default'; + const selectedPaths = selectedModelPaths[modelId] ?? []; + return ( +
+
+
+ {namesById.get(modelId) ?? modelId} +
+
+ {modelId} +
+
+ {paths.length > 0 ? ( + <> + { + if (!value) return; + setPathModes((current) => ({ + ...current, + [modelId]: value as ModelPathMode, + })); + }} + disabled={certify.isPending} + className='w-full justify-start' + > + Default + Choose paths + All paths + + {mode === 'default' && ( +

+ Uses the upstream provider's normal model routing. +

+ )} + {mode === 'selected' && ( + + + + + +
+ {paths.map((path) => { + const checked = selectedPaths.includes(path.path); + return ( + + ); + })} +
+
+
+ )} + {mode === 'all' && ( +

+ All {paths.length} paths will run in parallel. +

+ )} + + ) : ( +

+ Only the provider default route is available. +

+ )} +
+ ); + })} + {modelsNeedingPath.length > 0 && ( +

+ Choose at least one path for each model using “Choose paths”. +

+ )} +
+ +
+
+ setCheckCache(value === true)} + disabled={certify.isPending} + /> + +
+ +
+ + {currentModel && ( +
+ + Probing {namesById.get(currentModel.id) ?? currentModel.id} + {currentModel.pathCount > 1 + ? ` across ${currentModel.pathCount} paths in parallel` + : ''}{' '} + — model {currentModel.index} of {currentModel.total} +
+ )} + + {results.length > 0 && ( + result.resultKey).join('|')} + defaultValue={results[0].resultKey} + className='space-y-3' + > +
+ + {results.map((result) => { + const status = resultStatus(result); + const Icon = + status === 'error' ? XCircle : STATUS_STYLES[status].icon; + return ( + + + + {namesById.get(result.modelId) ?? result.modelId} ·{' '} + {result.pathLabel} + + + ); + })} + +
+ {results.map((result) => ( + + {result.report ? ( + + ) : ( +
+ {result.error ?? 'Certification failed'} +
+ )} +
+ ))} +
+ )} +
+
+ ); +} diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index f21e7b64..ba665da0 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -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; +export type CertificationRow = z.infer; +export type CertificationGoal = z.infer; +export type ProviderCertification = z.infer; + +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; @@ -107,6 +154,7 @@ export type AdminModelPricing = z.infer; export type AdminModelArchitecture = z.infer< typeof AdminModelArchitectureSchema >; +export type CertificationPath = z.infer; export type ProviderModels = z.infer; export interface AdminModelAsModel { @@ -317,6 +365,17 @@ export class AdminService { ); } + static async certifyProvider( + providerId: number, + body: CertifyProviderRequest = {} + ): Promise { + const data = await apiClient.post( + `/admin/api/upstream-providers/${providerId}/certify`, + body + ); + return ProviderCertificationSchema.parse(data); + } + static async getProviderModels(providerId: number): Promise { const data = await apiClient.get( `/admin/api/upstream-providers/${providerId}/models`