fix: resolve upstream certification review findings

This commit is contained in:
9qeklajc
2026-10-01 01:37:48 +02:00
parent 41b624c436
commit 943aa8f6fa
10 changed files with 913 additions and 53 deletions
+17 -16
View File
@@ -1236,9 +1236,7 @@ async def get_provider_models(provider_id: str) -> dict[str, object]:
]
path_result = await session.exec(
select(ModelPathRow).where(
ModelPathRow.upstream_provider_id == provider_pk
)
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]]] = {}
@@ -1251,12 +1249,11 @@ async def get_provider_models(provider_id: str) -> dict[str, object]:
}
)
from ..upstream.model_paths import public_model_id
from ..upstream.model_paths import exposed_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(), [])
paths = paths_by_public_id.get(exposed_model_id(model).lower(), [])
certification_paths[model.id] = paths
return {
@@ -1570,19 +1567,11 @@ async def certify_upstream_provider(
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,
@@ -1652,13 +1641,25 @@ async def certify_upstream_provider(
(
model
for model in upstream.get_cached_models()
if model.id == model_id
or model.forwarded_model_id == model_id
if model.id == model_id or model.forwarded_model_id == model_id
),
None,
)
if model_obj is not None:
break
if selected_path is not None:
from ..upstream.model_paths import exposed_model_id
selected_id = exposed_model_id(model_obj) if model_obj else model_id
if (
payload.model_id is None
or selected_id is None
or selected_id.lower() != selector.model_id.lower()
):
raise HTTPException(
status_code=400,
detail="Model path does not match the selected model",
)
if model_obj is None:
from ..upstream.certification import (
STATUS_WARN,
+29 -5
View File
@@ -192,8 +192,8 @@ async def probe_upstream(
) -> ProbeResult:
"""Call the upstream's ``/models`` and a one-token completion.
A transport failure on either call is recorded on the result rather
than raised: a dead upstream is a ``fail`` row, not a failed request.
Each HTTP call, including its body read, has an elapsed-time deadline.
A transport failure is a ``fail`` row, not a failed admin request.
"""
base = base_url.rstrip("/")
result = ProbeResult(
@@ -213,6 +213,7 @@ async def probe_upstream(
try:
started = time.monotonic()
try:
async with asyncio.timeout(timeout):
response = await client.get(result.models_url, headers=headers)
result.models_status = response.status_code
result.models_latency_ms = round((time.monotonic() - started) * 1000, 2)
@@ -246,6 +247,7 @@ async def probe_upstream(
"allow_fallbacks": False,
}
try:
async with asyncio.timeout(timeout):
response = await client.post(
result.chat_url, json=request_body, headers=headers
)
@@ -507,6 +509,9 @@ def _reported_usd_cost(payload: dict[str, Any]) -> float:
total = coerce_rate(cost_details.get("total_cost"))
if total is not None and total > 0:
return total
inference = coerce_rate(cost_details.get("upstream_inference_cost"))
if inference is not None and inference > 0 and usage.get("is_byok"):
return inference + (coerce_rate(usage.get("cost")) or 0.0)
for source in (usage, payload):
for field in ("total_cost", "cost"):
value = coerce_rate(source.get(field))
@@ -830,7 +835,11 @@ async def run_live_checks(
if not check_cache:
rows.extend(skipped_cache_rows("Skipped — cache checks disabled."))
elif probe.chat_payload is None:
elif (
probe.chat_payload is None
or probe.chat_status is None
or not 200 <= probe.chat_status < 300
):
rows.extend(
skipped_cache_rows("Skipped — the completion probe did not succeed.")
)
@@ -928,7 +937,14 @@ async def _resolve_sats_usd_price(override: float | None) -> float | None:
def _model_from_usd_pricing(
model_id: str, prompt_usd: float, completion_usd: float, sats_to_usd: float
model_id: str,
prompt_usd: float,
completion_usd: float,
sats_to_usd: float,
*,
provider_fee: float = 1.0,
cache_read_usd: float | None = None,
cache_write_usd: float | None = None,
) -> "Model":
"""A throwaway ``Model`` carrying just enough to exercise the cost engine."""
from ..payment.models import (
@@ -951,7 +967,12 @@ def _model_from_usd_pricing(
tokenizer="unknown",
instruct_type=None,
),
pricing=Pricing(prompt=prompt_usd, completion=completion_usd),
pricing=Pricing(
prompt=prompt_usd * provider_fee,
completion=completion_usd * provider_fee,
input_cache_read=(cache_read_usd or 0.0) * provider_fee,
input_cache_write=(cache_write_usd or 0.0) * provider_fee,
),
sats_pricing=None,
per_request_limits=None,
top_provider=None,
@@ -1042,6 +1063,9 @@ async def certify_upstream_url(
resolved_prompt or 0.0,
resolved_completion or 0.0,
sats_to_usd or 1.0,
provider_fee=provider_fee,
cache_read_usd=_as_price(entry.get("cache_read_input_token_cost")),
cache_write_usd=_as_price(entry.get("cache_creation_input_token_cost")),
)
rows = await run_live_checks(
base_url,
+20 -3
View File
@@ -13,6 +13,7 @@ with ``httpx`` and never enter the billing path.
from __future__ import annotations
import asyncio
import time
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
@@ -130,9 +131,11 @@ async def _post_completion(
url: str,
body: dict[str, Any],
headers: dict[str, str],
timeout: float,
) -> tuple[int | None, dict[str, Any] | None, str | None, float]:
started = time.monotonic()
try:
async with asyncio.timeout(timeout):
response = await client.post(url, json=body, headers=headers)
except Exception as exc: # noqa: BLE001 - transport failure is a row status
latency = round((time.monotonic() - started) * 1000, 2)
@@ -180,7 +183,8 @@ async def probe_cache(
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.
retry on HTTP 400/422, and the second call mirrors whichever format
succeeded. Each call's elapsed deadline includes the response body.
"""
base = base_url.rstrip("/")
result = CacheProbeResult(
@@ -200,14 +204,16 @@ async def probe_cache(
result.chat_url,
_request_body(model_id, prefix, "cache_control", endpoint_tag),
headers,
timeout,
)
if not _is_2xx(first[0]) and first[0] is not None:
if first[0] in (400, 422):
result.request_format = "plain"
first = await _post_completion(
client,
result.chat_url,
_request_body(model_id, prefix, "plain", endpoint_tag),
headers,
timeout,
)
_record(result, first)
if not _is_2xx(first[0]):
@@ -217,6 +223,7 @@ async def probe_cache(
result.chat_url,
_request_body(model_id, prefix, result.request_format, endpoint_tag),
headers,
timeout,
)
_record(result, second)
finally:
@@ -282,7 +289,17 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]:
evidence,
)
raw_keys = _raw_cache_keys(payload.get("usage"))
known_write_fields = {
"cache_creation_input_tokens",
"prompt_tokens_details.cache_creation_tokens",
"prompt_tokens_details.cache_write_tokens",
"input_tokens_details.cache_write_tokens",
}
raw_keys = [
key
for key in _raw_cache_keys(payload.get("usage"))
if key not in known_write_fields
]
if raw_keys:
evidence["unrecognised_cache_fields"] = raw_keys
return certification_row(
@@ -0,0 +1,130 @@
"""Exact certification paths retain the provider model's forwarded identity."""
import json
from typing import Any
from unittest.mock import patch
import pytest
import respx
from httpx import AsyncClient, Response
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ModelPathRow
from routstr.proxy import reinitialize_upstreams
from routstr.upstream.model_paths import encode_model_path
from tests.integration.test_certify_endpoint import (
_admin_headers,
_make_provider,
_model_row,
)
@pytest.mark.integration
@pytest.mark.asyncio
@pytest.mark.parametrize("forwarded", ["anthropic/claude-opus-4.6", "remote-id"])
@pytest.mark.parametrize("enabled", [True, False])
@respx.mock
async def test_certify_forwarded_alias_listed_path_succeeds(
integration_session: AsyncSession,
integration_client: AsyncClient,
forwarded: str,
enabled: bool,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from routstr import proxy
for name in ("_upstreams", "_provider_map", "_unique_models"):
monkeypatch.setattr(proxy, name, getattr(proxy, name).copy())
base_url = "https://certify-upstream.example/v1"
respx.get(f"{base_url}/models").mock(return_value=Response(200, json={"data": []}))
chat = respx.post(f"{base_url}/chat/completions").mock(
return_value=Response(
200,
json={
"model": forwarded,
"usage": {"prompt_tokens": 5, "completion_tokens": 1},
},
)
)
provider = await _make_provider(integration_session)
model = _model_row(provider.id, model_id="local-alias") # type: ignore[arg-type]
model.forwarded_model_id = forwarded
model.enabled = enabled
model_path = encode_model_path(base_url, forwarded, "endpoint")
integration_session.add(model)
integration_session.add(
ModelPathRow(
upstream_provider_id=provider.id,
model_id=forwarded,
path=model_path,
endpoint_tag="endpoint",
provider_slug="mock",
provider_type="generic",
)
)
other_id = f"other/{forwarded.rsplit('/', 1)[-1]}"
other_path = encode_model_path(base_url, other_id, "other-endpoint")
integration_session.add(
ModelPathRow(
upstream_provider_id=provider.id,
model_id=other_id,
path=other_path,
endpoint_tag="other-endpoint",
provider_slug="mock",
provider_type="generic",
)
)
await integration_session.commit()
with (
patch("routstr.payment.models.sats_usd_price", return_value=0.0005),
patch("routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005),
patch("routstr.payment.price.SATS_USD_PRICE", 0.0005),
):
await reinitialize_upstreams()
listed = await integration_client.get(
f"/admin/api/upstream-providers/{provider.id}/models",
headers=_admin_headers(),
)
assert listed.status_code == 200, listed.text
assert (
listed.json()["certification_paths"]["local-alias"][0]["path"] == model_path
)
response = await integration_client.post(
f"/admin/api/upstream-providers/{provider.id}/certify",
headers=_admin_headers(),
json={
"model_id": "local-alias",
"model_path": model_path,
"check_cache": False,
},
)
if not enabled:
assert response.status_code == 400, response.text
assert chat.call_count == 0
return
assert response.status_code == 200, response.text
assert chat.call_count == 1
body: dict[str, Any] = json.loads(chat.calls[0].request.content)
assert body["model"] == forwarded
assert body["provider"] == {"order": ["endpoint"], "allow_fallbacks": False}
mismatch = await integration_client.post(
f"/admin/api/upstream-providers/{provider.id}/certify",
headers=_admin_headers(),
json={
"model_id": "wrong-alias",
"model_path": model_path,
"check_cache": False,
},
)
assert mismatch.status_code == 400
wrong_prefix = await integration_client.post(
f"/admin/api/upstream-providers/{provider.id}/certify",
headers=_admin_headers(),
json={
"model_id": "local-alias",
"model_path": other_path,
"check_cache": False,
},
)
assert wrong_prefix.status_code == 400
assert chat.call_count == 1
@@ -0,0 +1,308 @@
"""Regression coverage for certification pricing and probe lifecycle fixes."""
import asyncio
from collections.abc import AsyncIterator
from typing import Any
import httpx
import pytest
from routstr.payment import price
from routstr.payment.cost_calculation import calculate_cost
from routstr.upstream.certification import (
_model_from_usd_pricing,
certify_upstream_url,
cost_prompt_completion_row,
probe_upstream,
run_live_checks,
)
from routstr.upstream.certification_cache import (
CacheProbeResult,
cache_reported_row,
cost_margin_row,
probe_cache,
)
@pytest.mark.asyncio
async def test_byok_cost_and_margin_include_inference_and_routing(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(price, "SATS_USD_PRICE", 0.001)
model = _model_from_usd_pricing("test-model", 1e-6, 2e-6, 0.001)
payload = {
"model": "test-model",
"usage": {
"prompt_tokens": 10,
"completion_tokens": 1,
"is_byok": True,
"cost": 0.00005,
"cost_details": {"upstream_inference_cost": 0.001},
},
}
from routstr.upstream.certification import ProbeResult
probe = ProbeResult(
base_url="https://mock.example/v1",
models_url="https://mock.example/v1/models",
chat_url="https://mock.example/v1/chat/completions",
chat_status=200,
chat_payload=payload,
)
cost = await calculate_cost(payload, 1_000_000, model_obj=model, provider_fee=1)
row = cost_prompt_completion_row(
model=model, probe=probe, cost_data=cost, provider_fee=1, sats_to_usd=0.001
)
assert row["status"] == "ok"
assert row["evidence"]["expected_total_msats"] == 1050
margin = cost_margin_row(
model=model, payloads=[payload], provider_fee=1, sats_to_usd=0.001
)
assert margin["status"] == "fail"
assert margin["evidence"]["samples"][0]["upstream_msats_with_fee"] == 1050
@pytest.mark.asyncio
@pytest.mark.parametrize(
("usage", "entry", "fee", "explicit", "expected"),
[
(
{"prompt_tokens": 1000, "completion_tokens": 1},
{"input_cost_per_token": 9e-6, "output_cost_per_token": 9e-6},
2,
True,
2004,
),
(
{
"prompt_tokens": 1000,
"completion_tokens": 1,
"prompt_tokens_details": {"cached_tokens": 900},
},
{
"cache_read_input_token_cost": -1,
"cache_creation_input_token_cost": float("inf"),
},
1,
False,
1002,
),
(
{
"prompt_tokens": 1000,
"completion_tokens": 1,
"prompt_tokens_details": {"cached_tokens": 900},
},
{"cache_read_input_token_cost": 1e-7},
1,
False,
192,
),
(
{
"input_tokens": 100,
"output_tokens": 1,
"cache_creation_input_tokens": 900,
},
{"cache_creation_input_token_cost": 1.25e-6},
2,
False,
2454,
),
(
{"prompt_tokens": 1000, "completion_tokens": 1, "cost": 0.001002},
{},
2,
True,
2004,
),
],
)
async def test_standalone_preserves_fee_cache_rates_and_usd_fee(
monkeypatch: pytest.MonkeyPatch,
usage: dict[str, Any],
entry: dict[str, float],
fee: float,
explicit: bool,
expected: int,
) -> None:
from routstr.payment import models
monkeypatch.setattr(price, "SATS_USD_PRICE", None)
monkeypatch.setattr(price, "BTC_USD_PRICE", None)
monkeypatch.setattr(
models,
"litellm_cost_entry",
lambda _: {
"input_cost_per_token": 1e-6,
"output_cost_per_token": 2e-6,
**entry,
},
)
def handle(request: httpx.Request) -> httpx.Response:
if request.url.path.endswith("/models"):
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
return httpx.Response(200, json={"model": "test-model", "usage": usage})
kwargs = {"prompt_price": 1e-6, "completion_price": 2e-6} if explicit else {}
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
result = await certify_upstream_url(
"https://mock.example/v1",
model_id="test-model",
provider_fee=fee,
sats_usd_price=0.001,
check_cache=False,
client=client,
**kwargs,
)
row = next(row for row in result["rows"] if row["id"] == "cost.prompt_completion")
assert row["status"] == "ok", row
assert row["evidence"]["actual_total_msats"] == expected
class TrickleBody(httpx.AsyncByteStream):
def __init__(self) -> None:
self.closed = False
self.started = asyncio.Event()
async def __aiter__(self) -> AsyncIterator[bytes]:
self.started.set()
for chunk in (b'{"data":', b"[]", b"}"):
await asyncio.sleep(0.03)
yield chunk
async def aclose(self) -> None:
self.closed = True
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ["models", "chat", "cache"])
async def test_probe_elapsed_deadline_closes_trickling_response(mode: str) -> None:
body = TrickleBody()
def handle(request: httpx.Request) -> httpx.Response:
if mode == "chat" and request.method == "GET":
return httpx.Response(200, json={"data": []})
return httpx.Response(200, stream=body)
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
if mode == "cache":
result = await probe_cache(
"https://mock.example/v1", "", "test-model", client=client, timeout=0.05
)
assert result.statuses == [None]
assert "TimeoutError" in (result.errors[0] or "")
else:
probe = await probe_upstream(
"https://mock.example/v1",
"",
"test-model" if mode == "chat" else "",
client=client,
timeout=0.05,
)
if mode == "chat":
assert probe.chat_status is None
assert "TimeoutError" in (probe.chat_error or "")
else:
assert probe.models_status is None
assert "TimeoutError" in (probe.models_error or "")
assert body.closed
assert not client.is_closed
@pytest.mark.asyncio
@pytest.mark.parametrize("cache", [False, True])
@pytest.mark.parametrize("owns_client", [False, True])
async def test_probe_cancellation_closes_body_and_owned_client(
monkeypatch: pytest.MonkeyPatch, cache: bool, owns_client: bool
) -> None:
body = TrickleBody()
client = httpx.AsyncClient(
transport=httpx.MockTransport(lambda _: httpx.Response(200, stream=body))
)
if owns_client:
monkeypatch.setattr(httpx, "AsyncClient", lambda **_: client)
probe = probe_cache if cache else probe_upstream
task = asyncio.create_task(
probe("https://mock.example/v1", "", "", client=None if owns_client else client)
)
await body.started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert body.closed
assert client.is_closed == owns_client
await client.aclose()
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [401, 403, 429, 500, 503])
async def test_failed_initial_completion_skips_cache(status: int) -> None:
calls = []
def handle(request: httpx.Request) -> httpx.Response:
if request.url.path.endswith("/models"):
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
calls.append(request)
return httpx.Response(status, json={"error": "failed"})
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
rows = await run_live_checks(
"https://mock.example/v1",
"",
_model_from_usd_pricing("test-model", 1e-6, 2e-6, 0.001),
provider_fee=1,
sats_to_usd=0.001,
client=client,
)
assert len(calls) == 1
assert (
next(row for row in rows if row["id"] == "cache.reported")["status"] == "warn"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("status", [401, 403, 429, 500, 503])
async def test_cache_does_not_retry_non_format_errors(status: int) -> None:
calls = []
def handle(request: httpx.Request) -> httpx.Response:
calls.append(request)
return httpx.Response(status, json={"error": "failed"})
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
result = await probe_cache(
"https://mock.example/v1", "", "test-model", client=client
)
assert len(calls) == 1
assert result.statuses == [status]
@pytest.mark.parametrize(
"usage",
[
{"cache_creation_input_tokens": 3000},
{"prompt_tokens_details": {"cache_creation_tokens": 3000}},
{"prompt_tokens_details": {"cache_write_tokens": 3000}},
{"input_tokens_details": {"cache_write_tokens": 3000}},
],
)
@pytest.mark.parametrize("unknown", [False, True])
def test_known_cache_writes_are_no_hit_not_unrecognized(
usage: dict[str, Any], unknown: bool
) -> None:
usage = {"input_tokens": 10, "output_tokens": 1, **usage}
if unknown:
usage["unknown_cached_read_tokens"] = 5
payload = {"usage": usage}
row = cache_reported_row(
CacheProbeResult(
chat_url="mock", statuses=[200, 200], payloads=[payload, payload]
)
)
assert row["status"] == ("fail" if unknown else "warn")
assert row["evidence"]["second_usage"]["cache_write_tokens"] == 3000
assert row["evidence"].get("unrecognised_cache_fields", []) == (
["unknown_cached_read_tokens"] if unknown else []
)
+32 -6
View File
@@ -1,6 +1,6 @@
'use client';
import { useEffect, useMemo, useState } from 'react';
import { useEffect, useMemo, useRef, useState } from 'react';
import Link from 'next/link';
import { useQueries, useQuery } from '@tanstack/react-query';
import {
@@ -58,6 +58,7 @@ import {
countCertificationTargets,
emptyCertificationSetup,
getModelsNeedingPath,
getSelectedCertificationResults,
} from '@/lib/provider-certification';
import type {
CertificationProgress,
@@ -86,6 +87,14 @@ export default function MultiProviderCertificationPage() {
Record<number, CertificationProgress | null>
>({});
const [runningProviderIds, setRunningProviderIds] = useState<number[]>([]);
const activeRuns = useRef(new Set<number>());
const generation = useRef(0);
useEffect(() => {
return () => {
generation.current += 1;
};
}, []);
const providersQuery = useQuery({
queryKey: ['upstream-providers'],
@@ -180,6 +189,10 @@ export default function MultiProviderCertificationPage() {
};
const runOneProvider = async (providerId: number) => {
if (activeRuns.current.has(providerId)) return;
activeRuns.current.add(providerId);
const runGeneration = generation.current;
const shouldContinue = () => generation.current === runGeneration;
const setup = setups[providerId] ?? emptyCertificationSetup();
const modelRuns = providerRuns(providerId);
setResultsByProvider((current) => ({ ...current, [providerId]: [] }));
@@ -192,18 +205,27 @@ export default function MultiProviderCertificationPage() {
providerId,
modelRuns,
includeCache: setup.checkCache,
onProgress: (progress) =>
shouldContinue,
onProgress: (progress) => {
if (shouldContinue()) {
setProgressByProvider((current) => ({
...current,
[providerId]: progress,
})),
onResults: (results) =>
}));
}
},
onResults: (results) => {
if (shouldContinue()) {
setResultsByProvider((current) => ({
...current,
[providerId]: results,
})),
}));
}
},
});
} finally {
activeRuns.current.delete(providerId);
if (shouldContinue()) {
setProgressByProvider((current) => ({
...current,
[providerId]: null,
@@ -212,6 +234,7 @@ export default function MultiProviderCertificationPage() {
current.filter((id) => id !== providerId)
);
}
}
};
const runAllProviders = () => {
@@ -241,7 +264,10 @@ export default function MultiProviderCertificationPage() {
);
const allReady =
selectedProviderIds.length > 0 && incompleteProviders.length === 0;
const allResults = Object.values(resultsByProvider).flat();
const allResults = getSelectedCertificationResults(
selectedProviderIds,
resultsByProvider
);
const aggregateSummary = summarizeCertificationResults(allResults);
const pendingRoutes = Math.max(totalRoutes - allResults.length, 0);
@@ -80,7 +80,13 @@ export function ProviderCertificationDialog({
};
return (
<Dialog open={open} onOpenChange={onOpenChange}>
<Dialog
open={open}
onOpenChange={(nextOpen) => {
if (!nextOpen) reset();
onOpenChange(nextOpen);
}}
>
<DialogContent className='flex h-[90dvh] max-h-[90dvh] flex-col overflow-hidden sm:max-w-[780px]'>
<DialogHeader className='shrink-0'>
<DialogTitle>Certify upstream models</DialogTitle>
@@ -147,7 +153,9 @@ export function ProviderCertificationDialog({
<RotateCcw className='h-4 w-4' />
)}
{isPending
? `Running ${progress?.modelIndex ?? 1} of ${progress?.modelTotal ?? setup.selectedModelIds.length}`
? progress
? `Running ${progress.modelIndex} of ${progress.modelTotal}`
: 'Finishing in-flight probe'
: results.length > 0
? `Run ${targetCount} route${targetCount === 1 ? '' : 's'} again`
: `Certify ${targetCount || ''} route${targetCount === 1 ? '' : 's'}`}
@@ -0,0 +1,328 @@
import assert from 'node:assert/strict';
import { readFileSync } from 'node:fs';
import { test } from 'node:test';
import { fileURLToPath } from 'node:url';
import vm from 'node:vm';
import ts from 'typescript';
function loadSource(path, imports) {
const source = readFileSync(new URL(path, import.meta.url), 'utf8');
const { outputText } = ts.transpileModule(source, {
compilerOptions: {
module: ts.ModuleKind.CommonJS,
target: ts.ScriptTarget.ES2020,
jsx: ts.JsxEmit.ReactJSX,
},
fileName: fileURLToPath(new URL(path, import.meta.url)),
});
const sourceModule = { exports: {} };
vm.runInNewContext(outputText, {
module: sourceModule,
exports: sourceModule.exports,
require(name) {
assert.ok(name in imports, `Unexpected import ${name}`);
return imports[name];
},
});
return sourceModule.exports;
}
function runnerHarness() {
const calls = [];
const pending = [];
const slots = [];
const cleanups = [];
let cursor = 0;
const react = {
useState(initial) {
const index = cursor++;
if (!(index in slots)) slots[index] = initial;
return [
slots[index],
(next) => {
slots[index] = next;
},
];
},
useRef(initial) {
const index = cursor++;
if (!(index in slots)) slots[index] = { current: initial };
return slots[index];
},
useEffect(effect) {
const index = cursor++;
if (!(index in slots)) {
slots[index] = true;
cleanups.push(effect());
}
},
useCallback: (fn) => fn,
};
const exports = loadSource('./use-provider-certification-runner.ts', {
react,
'@/lib/api/services/admin': {
AdminService: {
certifyProvider(providerId, options) {
calls.push({ providerId, ...options });
return new Promise((resolve, reject) =>
pending.push({ resolve, reject })
);
},
},
},
'@/lib/provider-certification': {
getErrorMessage: (error) => error.message,
},
});
return {
...exports,
calls,
pending,
render() {
cursor = 0;
return exports.useProviderCertificationRunner(1);
},
unmount() {
cleanups.forEach((cleanup) => cleanup?.());
},
};
}
const model = (modelId, paths = ['default']) => ({
modelId,
targets: paths.map((path) => ({ path, label: path })),
});
const report = { rows: [] };
const flush = () => new Promise((resolve) => setImmediate(resolve));
test('reset cancels queued paths/models and blocks restart until in-flight work finishes', async () => {
const harness = runnerHarness();
const hook = harness.render();
const old = hook.run([model('first', ['a', 'b']), model('second')], false);
assert.equal(harness.render().isPending, true);
hook.reset();
assert.equal(harness.render().isPending, true);
await hook.run([model('restart')], false);
assert.equal(harness.calls.length, 1);
harness.pending[0].resolve(report);
await old;
const finished = harness.render();
assert.equal(finished.isPending, false);
assert.equal(finished.progress, null);
assert.equal(finished.results.length, 0);
assert.deepEqual(
harness.calls.map((call) => call.model_id),
['first']
);
const fresh = finished.run([model('restart')], false);
harness.pending[1].resolve(report);
await fresh;
assert.deepEqual(
harness.calls.map((call) => call.model_id),
['first', 'restart']
);
});
test('unmount cancels queued requests and suppresses stale result updates', async () => {
const harness = runnerHarness();
const old = harness.render().run([model('first'), model('second')], false);
harness.unmount();
harness.pending[0].resolve(report);
await old;
assert.equal(harness.calls.length, 1);
assert.equal(harness.render().results.length, 0);
});
test('cancellation before dispatch makes no request or progress callback', async () => {
const harness = runnerHarness();
const progress = [];
const result = await harness.runProviderCertification({
providerId: 1,
modelRuns: [model('first')],
includeCache: false,
shouldContinue: () => false,
onProgress: (next) => progress.push(next),
});
assert.equal(result.length, 0);
assert.equal(harness.calls.length, 0);
assert.equal(progress.length, 0);
});
test('per-route failures remain isolated and normal runs retain completed results', async () => {
const harness = runnerHarness();
const done = harness.render().run([model('first', ['a', 'b'])], true);
harness.pending[0].reject(new Error('route failed'));
await flush();
harness.pending[1].resolve(report);
await done;
const finished = harness.render();
assert.equal(finished.isPending, false);
assert.equal(finished.results.length, 2);
assert.equal(finished.results[0].error, 'route failed');
assert.equal(finished.results[1].report, report);
assert.equal(harness.calls[1].check_cache, true);
});
test('dialog close cancels synchronously before notifying its owner', () => {
const events = [];
const jsx = (type, props) => ({ type, props });
const components = new Proxy({}, { get: (_, name) => name });
const imports = {
react: {
useState: (initial) => [
typeof initial === 'function' ? initial() : initial,
() => {},
],
useEffect: (effect) => effect(),
},
'react/jsx-runtime': { jsx, jsxs: jsx },
'@tanstack/react-query': { useQuery: () => ({}) },
'lucide-react': components,
'@/components/ui/dialog': components,
'@/components/ui/button': components,
'@/components/ui/badge': components,
'@/components/ui/tabs': components,
'@/components/provider-certification-results': components,
'@/components/provider-certification-setup': {
...components,
ProviderCertificationSetupPanel: 'SetupPanel',
getCertificationModelNames: () => ({}),
},
'@/hooks/use-provider-certification-runner': {
useProviderCertificationRunner: () => ({
results: [],
progress: null,
isPending: false,
run: () => {},
reset: () => events.push('reset'),
}),
},
'@/lib/api/services/admin': { AdminService: {} },
'@/lib/provider-certification': {
emptyCertificationSetup: () => ({
selectedModelIds: [],
checkCache: false,
}),
buildModelRuns: () => [],
countCertificationTargets: () => 0,
getModelsNeedingPath: () => [],
},
};
const { ProviderCertificationDialog } = loadSource(
'../components/provider-certification-dialog.tsx',
imports
);
const dialog = ProviderCertificationDialog({
open: true,
provider: { id: 1, provider_type: 'generic' },
onOpenChange: () => events.push('owner'),
});
dialog.props.onOpenChange(false);
assert.deepEqual(events, ['reset', 'owner']);
});
test('multi-provider page guards duplicate starts and cancels queued work on unmount', async () => {
const harness = runnerHarness();
const cleanups = [];
const updates = [];
let stateIndex = 0;
const setup = {
selectedModelIds: ['first', 'second'],
pathModes: {},
selectedModelPaths: {},
checkCache: false,
};
const initialStates = [[1], 1, { 1: setup }];
const react = {
useState(initial) {
const index = stateIndex++;
return [
index < 3 ? initialStates[index] : initial,
(next) => updates.push(next),
];
},
useMemo: (fn) => fn(),
useRef: (initial) => ({ current: initial }),
useEffect: (effect) => cleanups.push(effect()),
};
const helpers = loadSource('../lib/provider-certification.ts', {
'@/lib/api/errors': { getApiErrorMessage: () => '' },
});
const jsx = (type, props) => ({ type, props });
const components = new Proxy({}, { get: (_, name) => name });
const imports = {
react,
'react/jsx-runtime': { jsx, jsxs: jsx },
'next/link': { default: 'Link' },
'@tanstack/react-query': {
useQuery: () => ({ data: [{ id: 1, provider_type: 'generic' }] }),
useQueries: () => [{ data: { certification_paths: {} } }],
},
'lucide-react': components,
'@/components/provider-certification-results': {
summarizeCertificationResults: () => ({}),
ProviderCertificationResults: 'Results',
},
'@/components/provider-certification-setup': {
getCertificationModelNames: () => ({}),
ProviderCertificationSetupPanel: 'Setup',
},
'@/hooks/use-provider-certification-runner': harness,
'@/lib/api/services/admin': { AdminService: {} },
'@/lib/provider-certification': helpers,
'@/lib/utils': { cn: () => '' },
};
for (const name of [
'app-page-shell',
'page-header',
'ui/badge',
'ui/button',
'ui/card',
'ui/checkbox',
'ui/command',
'ui/popover',
'ui/select',
'ui/tabs',
])
imports[`@/components/${name}`] = components;
const { default: Page } = loadSource(
'../app/providers/certification/page.tsx',
imports
);
const nodes = [];
const visit = (node) => {
if (!node || typeof node !== 'object') return;
if (Array.isArray(node)) return node.forEach(visit);
nodes.push(node);
visit(node.props?.children);
};
visit(Page());
const runAll = nodes.find(
(node) => node.props?.onClick?.name === 'runAllProviders'
);
assert.ok(runAll);
runAll.props.onClick();
runAll.props.onClick();
assert.equal(harness.calls.length, 1);
cleanups.forEach((cleanup) => cleanup?.());
const beforeCompletion = updates.length;
harness.pending[0].resolve(report);
await flush();
assert.equal(harness.calls.length, 1);
assert.equal(updates.length, beforeCompletion);
});
test('selected-provider aggregate excludes deselected providers and restores them on reselection', () => {
const { getSelectedCertificationResults } = loadSource(
'../lib/provider-certification.ts',
{
'@/lib/api/errors': { getApiErrorMessage: () => '' },
}
);
const a = { providerId: 1, resultKey: 'a' };
const results = { 1: [a] };
const selected = getSelectedCertificationResults([2], results);
assert.equal(selected.length, 0);
assert.equal(Math.max(1 - selected.length, 0), 1);
assert.deepEqual([...getSelectedCertificationResults([1, 2], results)], [a]);
});
+15 -3
View File
@@ -14,6 +14,7 @@ interface RunProviderCertificationOptions {
providerId: number;
modelRuns: ModelRun[];
includeCache: boolean;
shouldContinue?: () => boolean;
onProgress?: (progress: CertificationProgress | null) => void;
onResults?: (results: ModelCertificationResult[]) => void;
}
@@ -22,12 +23,14 @@ export async function runProviderCertification({
providerId,
modelRuns,
includeCache,
shouldContinue = () => true,
onProgress,
onResults,
}: RunProviderCertificationOptions): Promise<ModelCertificationResult[]> {
const completed: ModelCertificationResult[] = [];
for (const [index, run] of modelRuns.entries()) {
if (!shouldContinue()) break;
onProgress?.({
modelId: run.modelId,
modelIndex: index + 1,
@@ -38,6 +41,7 @@ export async function runProviderCertification({
// parallel paths multiply that spend and the admin request load.
const batch: ModelCertificationResult[] = [];
for (const target of run.targets) {
if (!shouldContinue()) break;
const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`;
try {
const report = await AdminService.certifyProvider(providerId, {
@@ -62,11 +66,12 @@ export async function runProviderCertification({
});
}
}
if (!shouldContinue()) break;
completed.push(...batch);
onResults?.([...completed]);
}
onProgress?.(null);
if (shouldContinue()) onProgress?.(null);
return completed;
}
@@ -76,6 +81,7 @@ export function useProviderCertificationRunner(providerId: number) {
const [isPending, setIsPending] = useState(false);
const generation = useRef(0);
const mounted = useRef(true);
const active = useRef(false);
useEffect(() => {
mounted.current = true;
@@ -89,11 +95,14 @@ export function useProviderCertificationRunner(providerId: number) {
generation.current += 1;
setResults([]);
setProgress(null);
setIsPending(false);
// Already dispatched probes may still spend credit; keep them pending.
setIsPending(active.current);
}, []);
const run = useCallback(
async (modelRuns: ModelRun[], includeCache: boolean) => {
if (active.current || !mounted.current) return [];
active.current = true;
const runGeneration = generation.current + 1;
generation.current = runGeneration;
setResults([]);
@@ -104,6 +113,8 @@ export function useProviderCertificationRunner(providerId: number) {
providerId,
modelRuns,
includeCache,
shouldContinue: () =>
mounted.current && generation.current === runGeneration,
onProgress: (nextProgress) => {
if (mounted.current && generation.current === runGeneration) {
setProgress(nextProgress);
@@ -116,7 +127,8 @@ export function useProviderCertificationRunner(providerId: number) {
},
});
} finally {
if (mounted.current && generation.current === runGeneration) {
active.current = false;
if (mounted.current) {
setProgress(null);
setIsPending(false);
}
+6
View File
@@ -104,6 +104,12 @@ export const buildModelRuns = (
export const countCertificationTargets = (runs: ModelRun[]): number =>
runs.reduce((total, run) => total + run.targets.length, 0);
export const getSelectedCertificationResults = (
providerIds: number[],
resultsByProvider: Record<number, ModelCertificationResult[]>
): ModelCertificationResult[] =>
providerIds.flatMap((providerId) => resultsByProvider[providerId] ?? []);
export const getCertificationResultStatus = (
result: ModelCertificationResult
): CertificationStatus | 'error' => {