mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: resolve upstream certification review findings
This commit is contained in:
+17
-16
@@ -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,
|
||||
|
||||
@@ -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,7 +213,8 @@ async def probe_upstream(
|
||||
try:
|
||||
started = time.monotonic()
|
||||
try:
|
||||
response = await client.get(result.models_url, headers=headers)
|
||||
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)
|
||||
try:
|
||||
@@ -246,9 +247,10 @@ async def probe_upstream(
|
||||
"allow_fallbacks": False,
|
||||
}
|
||||
try:
|
||||
response = await client.post(
|
||||
result.chat_url, json=request_body, headers=headers
|
||||
)
|
||||
async with asyncio.timeout(timeout):
|
||||
response = await client.post(
|
||||
result.chat_url, json=request_body, headers=headers
|
||||
)
|
||||
result.chat_status = response.status_code
|
||||
result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2)
|
||||
try:
|
||||
@@ -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,
|
||||
|
||||
@@ -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,10 +131,12 @@ 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:
|
||||
response = await client.post(url, json=body, headers=headers)
|
||||
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)
|
||||
return None, None, f"{type(exc).__name__}: {exc}", latency
|
||||
@@ -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 []
|
||||
)
|
||||
@@ -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,25 +205,35 @@ export default function MultiProviderCertificationPage() {
|
||||
providerId,
|
||||
modelRuns,
|
||||
includeCache: setup.checkCache,
|
||||
onProgress: (progress) =>
|
||||
setProgressByProvider((current) => ({
|
||||
...current,
|
||||
[providerId]: progress,
|
||||
})),
|
||||
onResults: (results) =>
|
||||
setResultsByProvider((current) => ({
|
||||
...current,
|
||||
[providerId]: results,
|
||||
})),
|
||||
shouldContinue,
|
||||
onProgress: (progress) => {
|
||||
if (shouldContinue()) {
|
||||
setProgressByProvider((current) => ({
|
||||
...current,
|
||||
[providerId]: progress,
|
||||
}));
|
||||
}
|
||||
},
|
||||
onResults: (results) => {
|
||||
if (shouldContinue()) {
|
||||
setResultsByProvider((current) => ({
|
||||
...current,
|
||||
[providerId]: results,
|
||||
}));
|
||||
}
|
||||
},
|
||||
});
|
||||
} finally {
|
||||
setProgressByProvider((current) => ({
|
||||
...current,
|
||||
[providerId]: null,
|
||||
}));
|
||||
setRunningProviderIds((current) =>
|
||||
current.filter((id) => id !== providerId)
|
||||
);
|
||||
activeRuns.current.delete(providerId);
|
||||
if (shouldContinue()) {
|
||||
setProgressByProvider((current) => ({
|
||||
...current,
|
||||
[providerId]: null,
|
||||
}));
|
||||
setRunningProviderIds((current) =>
|
||||
current.filter((id) => id !== providerId)
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
@@ -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]);
|
||||
});
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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' => {
|
||||
|
||||
Reference in New Issue
Block a user