From 943aa8f6fa7e5c1f038682538a3083ad4dbfe9a8 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 01:37:48 +0200 Subject: [PATCH] fix: resolve upstream certification review findings --- routstr/core/admin.py | 33 +- routstr/upstream/certification.py | 42 ++- routstr/upstream/certification_cache.py | 25 +- tests/integration/test_certify_alias_paths.py | 130 +++++++ .../test_certification_review_regressions.py | 308 ++++++++++++++++ ui/app/providers/certification/page.tsx | 64 +++- .../provider-certification-dialog.tsx | 12 +- ...use-provider-certification-runner.test.mjs | 328 ++++++++++++++++++ ui/hooks/use-provider-certification-runner.ts | 18 +- ui/lib/provider-certification.ts | 6 + 10 files changed, 913 insertions(+), 53 deletions(-) create mode 100644 tests/integration/test_certify_alias_paths.py create mode 100644 tests/unit/test_certification_review_regressions.py create mode 100644 ui/hooks/use-provider-certification-runner.test.mjs diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 745bb8e4..bf3aac15 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -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, diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 188bd16b..1051b019 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -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, diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index d8f21f13..d8ed2dad 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -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( diff --git a/tests/integration/test_certify_alias_paths.py b/tests/integration/test_certify_alias_paths.py new file mode 100644 index 00000000..44e10f7c --- /dev/null +++ b/tests/integration/test_certify_alias_paths.py @@ -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 diff --git a/tests/unit/test_certification_review_regressions.py b/tests/unit/test_certification_review_regressions.py new file mode 100644 index 00000000..81dac641 --- /dev/null +++ b/tests/unit/test_certification_review_regressions.py @@ -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 [] + ) diff --git a/ui/app/providers/certification/page.tsx b/ui/app/providers/certification/page.tsx index ef1b0ef4..93274331 100644 --- a/ui/app/providers/certification/page.tsx +++ b/ui/app/providers/certification/page.tsx @@ -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 >({}); const [runningProviderIds, setRunningProviderIds] = useState([]); + const activeRuns = useRef(new Set()); + 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); diff --git a/ui/components/provider-certification-dialog.tsx b/ui/components/provider-certification-dialog.tsx index 381cda4e..ee8dd24b 100644 --- a/ui/components/provider-certification-dialog.tsx +++ b/ui/components/provider-certification-dialog.tsx @@ -80,7 +80,13 @@ export function ProviderCertificationDialog({ }; return ( - + { + if (!nextOpen) reset(); + onOpenChange(nextOpen); + }} + > Certify upstream models @@ -147,7 +153,9 @@ export function ProviderCertificationDialog({ )} {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'}`} diff --git a/ui/hooks/use-provider-certification-runner.test.mjs b/ui/hooks/use-provider-certification-runner.test.mjs new file mode 100644 index 00000000..f2cfbeb1 --- /dev/null +++ b/ui/hooks/use-provider-certification-runner.test.mjs @@ -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]); +}); diff --git a/ui/hooks/use-provider-certification-runner.ts b/ui/hooks/use-provider-certification-runner.ts index fb7a9f7d..8bd117aa 100644 --- a/ui/hooks/use-provider-certification-runner.ts +++ b/ui/hooks/use-provider-certification-runner.ts @@ -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 { 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); } diff --git a/ui/lib/provider-certification.ts b/ui/lib/provider-certification.ts index 1b9a0d20..dff4478c 100644 --- a/ui/lib/provider-certification.ts +++ b/ui/lib/provider-certification.ts @@ -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 +): ModelCertificationResult[] => + providerIds.flatMap((providerId) => resultsByProvider[providerId] ?? []); + export const getCertificationResultStatus = ( result: ModelCertificationResult ): CertificationStatus | 'error' => {