diff --git a/routstr/core/admin.py b/routstr/core/admin.py index bd5b4e18..745bb8e4 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1549,7 +1549,6 @@ async def certify_upstream_provider( :mod:`routstr.upstream.certification`, and a ``checklist`` of the operator-facing goals. """ - from ..payment.price import sats_usd_price from ..upstream.certification import ( MAX_PROBE_TIMEOUT_SECONDS, PROBE_TIMEOUT_SECONDS, @@ -1707,7 +1706,16 @@ async def certify_upstream_provider( *skipped_cache_rows("Skipped — no model to probe."), ] else: - sats_to_usd = sats_usd_price() + from ..payment import price as price_module + + # Never fetch the price inline: the lifespan task owns it, and a fetch + # here could block the request for the exchange timeout. + sats_to_usd = price_module.SATS_USD_PRICE + if not sats_to_usd: + raise HTTPException( + status_code=503, + detail="sats/USD price is not initialized yet; retry shortly", + ) if selected_path is not None: from ..upstream.model_paths import apply_model_path_pricing @@ -1717,8 +1725,8 @@ async def certify_upstream_provider( provider.provider_fee, sats_to_usd, ) - # Clamp the admin-supplied timeout so a probe cannot hold the request - # open indefinitely. + # Clamp the admin-supplied timeout per upstream call. The run makes up + # to five calls, so the request can stay open for up to five times it. requested = ( payload.timeout_seconds if payload.timeout_seconds is not None diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 5a87c8d0..188bd16b 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -515,6 +515,42 @@ def _reported_usd_cost(payload: dict[str, Any]) -> float: return 0.0 +def _fixed_token_pricing_active() -> bool: + """Whether node-wide fixed per-1k pricing overrides the model's rates.""" + from ..core.settings import settings + + return bool( + settings.fixed_pricing + and (settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens) + ) + + +def _token_rates(sats_pricing: Any) -> tuple[float, float, float, float]: + """The msats-per-1k rates the engine bills tokens at. + + Mirrors ``_get_pricing_rates``'s selection: node-wide fixed pricing + overrides the model's own rates, with cache tokens at the input rate. + + Returns ``(input, output, cache_read, cache_write)``. + """ + from ..core.settings import settings + + if _fixed_token_pricing_active(): + fixed_input = float(settings.fixed_per_1k_input_tokens) * 1000.0 + fixed_output = float(settings.fixed_per_1k_output_tokens) * 1000.0 + return fixed_input, fixed_output, fixed_input, fixed_input + + input_rate = float(sats_pricing.prompt) * 1_000_000.0 + output_rate = float(sats_pricing.completion) * 1_000_000.0 + cache_read_rate = ( + float(sats_pricing.input_cache_read or 0.0) * 1_000_000.0 or input_rate + ) + cache_write_rate = ( + float(sats_pricing.input_cache_write or 0.0) * 1_000_000.0 or input_rate + ) + return input_rate, output_rate, cache_read_rate, cache_write_rate + + def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int]: """Re-derive the token-priced charge independently of the engine. @@ -525,18 +561,10 @@ def _expected_token_msats(sats_pricing: Any, usage: Any) -> tuple[int, int, int] Returns ``(total_msats, input_msats, output_msats)``. Raises ``ValueError`` on a non-finite rate, which would otherwise crash ``math.ceil`` downstream. """ - input_rate = float(sats_pricing.prompt) * 1_000_000.0 - output_rate = float(sats_pricing.completion) * 1_000_000.0 - cache_read_rate = ( - float(sats_pricing.input_cache_read or 0.0) * 1_000_000.0 or input_rate - ) - cache_write_rate = ( - float(sats_pricing.input_cache_write or 0.0) * 1_000_000.0 or input_rate - ) - - rates = (input_rate, output_rate, cache_read_rate, cache_write_rate) + rates = _token_rates(sats_pricing) if not all(math.isfinite(rate) for rate in rates): raise ValueError(f"non-finite pricing rate in {rates!r}") + input_rate, output_rate, cache_read_rate, cache_write_rate = rates calc_input = round(usage.input_tokens / 1000 * input_rate, 3) calc_output = round(usage.output_tokens / 1000 * output_rate, 3) @@ -587,6 +615,18 @@ def cost_prompt_completion_row( "sats_usd_price": sats_to_usd, } + # Checked before the engine's error: with no price the engine cannot + # succeed, and that is a gap in the run's inputs, not a node fault. + if not pricing_known: + return certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + "No pricing is known for this model, so the charge cannot be " + "verified. Configure the model on the node, or pass explicit " + "prices, to certify this row.", + evidence, + ) if isinstance(cost_data, CostDataError): evidence["error"] = cost_data.message return certification_row( @@ -613,16 +653,6 @@ def cost_prompt_completion_row( "verify the charge against.", evidence, ) - if not pricing_known: - return certification_row( - "cost.prompt_completion", - STATUS_WARN, - "Prompt and completion cost calculated", - "No pricing is known for this model, so the charge cannot be " - "verified. Configure the model on the node, or pass explicit " - "prices, to certify this row.", - evidence, - ) reported_usd = _reported_usd_cost(payload) try: @@ -865,15 +895,24 @@ async def _resolve_sats_usd_price(override: float | None) -> float | None: feed once, and return ``None`` rather than raising so the cost row can degrade to a ``warn`` and the rest of the report still prints. """ - if override is not None: - return override if math.isfinite(override) and override > 0 else None - from ..payment import price as price_module + # The cost engine reads the module globals rather than this return value, + # so a resolved price is published there too or every token-priced cost + # row fails on "SATS price not initialized". + if override is not None: + if not (math.isfinite(override) and override > 0): + return None + price_module.SATS_USD_PRICE = override + price_module.BTC_USD_PRICE = override * price_module.SATS_PER_BTC + return override + if price_module.SATS_USD_PRICE: return float(price_module.SATS_USD_PRICE) if price_module.BTC_USD_PRICE: - return float(price_module.BTC_USD_PRICE) / price_module.SATS_PER_BTC + sats_price = float(price_module.BTC_USD_PRICE) / price_module.SATS_PER_BTC + price_module.SATS_USD_PRICE = sats_price + return sats_price try: await price_module._update_prices() diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index 3ee74e7a..d8f21f13 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -32,7 +32,9 @@ from .certification import ( STATUS_WARN, _expected_token_msats, _expected_usd_msats, + _fixed_token_pricing_active, _reported_usd_cost, + _token_rates, certification_row, safe_row, ) @@ -341,8 +343,7 @@ def cache_billing_row( ) pricing = model.sats_pricing - cache_read_rate = float(pricing.input_cache_read or 0.0) - input_rate = float(pricing.prompt) + input_rate, _, cache_read_rate, _ = _token_rates(pricing) full_usage = NormalizedUsage( input_tokens=usage.input_tokens + usage.cache_read_tokens @@ -367,8 +368,8 @@ def cache_billing_row( evidence.update( { "usage": usage.dict(), - "cache_read_rate_sats": cache_read_rate, - "input_rate_sats": input_rate, + "cache_read_rate_msats_per_1k": cache_read_rate, + "input_rate_msats_per_1k": input_rate, "actual_total_msats": actual_total, "expected_total_msats": expected_total, "full_price_total_msats": full_total, @@ -396,13 +397,18 @@ def cache_billing_row( evidence, ) if cache_read_rate <= 0.0 or cache_read_rate >= input_rate: + reason = ( + "the node uses fixed per-1k pricing" + if _fixed_token_pricing_active() + else "no discounted cache-read rate is configured" + ) return certification_row( ROW_BILLING, STATUS_WARN, TITLE_BILLING, f"Cached reads are billed at the full input rate ({actual_total} " - "msats) because no discounted cache-read rate is configured; " - "clients pay more than the upstream charges.", + f"msats) because {reason}; clients pay more than the upstream " + "charges.", evidence, ) return certification_row( diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index c7aa9151..ac38a9e0 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -399,7 +399,7 @@ async def test_certify_uses_selected_path_pricing_for_margin( respx.post(f"{base_url}/chat/completions").mock(side_effect=_respond) sats_usd = 0.0008616302499999999 - with patch("routstr.payment.price.sats_usd_price", return_value=sats_usd): + with patch("routstr.payment.price.SATS_USD_PRICE", sats_usd): resp = await integration_client.post( f"/admin/api/upstream-providers/{provider_id}/certify", headers=_admin_headers(), @@ -416,6 +416,24 @@ async def test_certify_uses_selected_path_pricing_for_margin( assert "15 < 26" in margin["detail"] +@pytest.mark.integration +@pytest.mark.asyncio +async def test_certify_returns_503_when_price_uninitialized( + integration_client: AsyncClient, integration_session: AsyncSession +) -> None: + provider_id = await _seed_and_init(integration_session, integration_client) + + with patch("routstr.payment.price.SATS_USD_PRICE", None): + resp = await integration_client.post( + f"/admin/api/upstream-providers/{provider_id}/certify", + headers=_admin_headers(), + json={}, + ) + + assert resp.status_code == 503, resp.text + assert "sats/USD price is not initialized" in resp.json()["detail"] + + @pytest.mark.integration @pytest.mark.asyncio @respx.mock diff --git a/tests/unit/test_certification.py b/tests/unit/test_certification.py index 1e963a79..8c1d82e4 100644 --- a/tests/unit/test_certification.py +++ b/tests/unit/test_certification.py @@ -10,6 +10,8 @@ from __future__ import annotations from typing import Any +import pytest + from routstr.upstream.certification import ( STATUS_FAIL, STATUS_OK, @@ -606,3 +608,122 @@ class TestBuildChecklist: goals = {item["goal"]: item["status"] for item in checklist} assert goals["heartbeat"] == STATUS_FAIL assert goals["usage_data"] == STATUS_WARN + + +# --- engine-consistent pricing ---------------------------------------------- + + +class TestStandaloneCostRow: + """The standalone runner prices a real completion through the engine.""" + + @staticmethod + def _client() -> Any: + import httpx + + def handler(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "cert-model"}]}) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-1", + "object": "chat.completion", + "model": "cert-model", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + }, + ) + + return httpx.AsyncClient(transport=httpx.MockTransport(handler)) + + @staticmethod + def _cost_row(result: dict[str, Any]) -> dict[str, Any]: + return next(r for r in result["rows"] if r["id"] == "cost.prompt_completion") + + async def test_override_price_reaches_the_engine( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.payment import price as price_module + from routstr.upstream.certification import certify_upstream_url + + monkeypatch.setattr(price_module, "SATS_USD_PRICE", None) + monkeypatch.setattr(price_module, "BTC_USD_PRICE", None) + + async with self._client() as client: + result = await certify_upstream_url( + "https://upstream.example/v1", + model_id="cert-model", + sats_usd_price=5e-7, + prompt_price=1e-6, + completion_price=2e-6, + client=client, + check_cache=False, + ) + + row = self._cost_row(result) + assert row["status"] == STATUS_OK, row["detail"] + + async def test_no_price_available_warns( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.payment import price as price_module + from routstr.upstream.certification import certify_upstream_url + + async def offline() -> None: + raise RuntimeError("exchange feed unreachable") + + monkeypatch.setattr(price_module, "SATS_USD_PRICE", None) + monkeypatch.setattr(price_module, "BTC_USD_PRICE", None) + monkeypatch.setattr(price_module, "_update_prices", offline) + + async with self._client() as client: + result = await certify_upstream_url( + "https://upstream.example/v1", + model_id="cert-model", + prompt_price=1e-6, + completion_price=2e-6, + client=client, + check_cache=False, + ) + + row = self._cost_row(result) + assert row["status"] == STATUS_WARN, row["detail"] + + +class TestFixedPricingCostRow: + async def test_fixed_pricing_node_certifies_ok( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.core.settings import settings + from routstr.payment import price as price_module + from routstr.payment.cost_calculation import calculate_cost + + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 3) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 7) + monkeypatch.setattr(price_module, "SATS_USD_PRICE", 0.0005) + + model = TestCostPromptCompletion()._model() + payload = { + "model": "test-model", + "usage": {"prompt_tokens": 1000, "completion_tokens": 500}, + } + cost_data = await calculate_cost( + payload, 1_000_000, model_obj=model, provider_fee=1.0 + ) + + row = cost_prompt_completion_row( + model=model, + probe=_probe(chat_status=200, chat_payload=payload), + cost_data=cost_data, + provider_fee=1.0, + sats_to_usd=0.0005, + ) + assert row["status"] == STATUS_OK, row["detail"] + assert row["evidence"]["expected_total_msats"] == 3000 + 3500 diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py index 111659ba..2d87a8cf 100644 --- a/tests/unit/test_certification_cache.py +++ b/tests/unit/test_certification_cache.py @@ -463,6 +463,32 @@ class TestErrorBranches: assert row["status"] == STATUS_FAIL assert "error" in row["evidence"] + @pytest.mark.asyncio + async def test_billing_warns_under_fixed_pricing( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + from routstr.core.settings import settings + from routstr.payment.cost_calculation import calculate_cost + + fixed_in, fixed_out = 2.0, 3.0 + monkeypatch.setattr(settings, "fixed_pricing", True) + monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", fixed_in) + monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", fixed_out) + monkeypatch.setattr("routstr.payment.price.SATS_USD_PRICE", SATS_USD) + + model = _model(cache_read=1.4e-8) + payload = _payload(CACHED) + cost = await calculate_cost(payload, 10**9, model_obj=model, provider_fee=1.0) + row = cache_billing_row( + model=model, + probe=_probe([_payload(UNCACHED), payload]), + cost_data=cost, + ) + assert row["status"] == STATUS_WARN, row + assert "fixed per-1k pricing" in row["detail"] + assert row["evidence"]["input_rate_msats_per_1k"] == fixed_in * 1000 + assert row["evidence"]["cache_read_rate_msats_per_1k"] == fixed_in * 1000 + def test_margin_fails_on_zero_sats_price(self) -> None: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) row = cost_margin_row( diff --git a/ui/hooks/use-provider-certification-runner.ts b/ui/hooks/use-provider-certification-runner.ts index aff3b914..fb7a9f7d 100644 --- a/ui/hooks/use-provider-certification-runner.ts +++ b/ui/hooks/use-provider-certification-runner.ts @@ -34,33 +34,34 @@ export async function runProviderCertification({ modelTotal: modelRuns.length, pathCount: run.targets.length, }); - const batch = await Promise.all( - run.targets.map(async (target): Promise => { - const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`; - try { - const report = await AdminService.certifyProvider(providerId, { - model_id: run.modelId, - model_path: target.path, - check_cache: includeCache, - }); - return { - resultKey, - providerId, - modelId: run.modelId, - pathLabel: target.label, - report, - }; - } catch (error) { - return { - resultKey, - providerId, - modelId: run.modelId, - pathLabel: target.label, - error: getErrorMessage(error), - }; - } - }) - ); + // Sequential on purpose: each run spends real upstream credits, and + // parallel paths multiply that spend and the admin request load. + const batch: ModelCertificationResult[] = []; + for (const target of run.targets) { + const resultKey = `${providerId}::${run.modelId}::${target.path ?? 'default'}`; + try { + const report = await AdminService.certifyProvider(providerId, { + model_id: run.modelId, + model_path: target.path, + check_cache: includeCache, + }); + batch.push({ + resultKey, + providerId, + modelId: run.modelId, + pathLabel: target.label, + report, + }); + } catch (error) { + batch.push({ + resultKey, + providerId, + modelId: run.modelId, + pathLabel: target.label, + error: getErrorMessage(error), + }); + } + } completed.push(...batch); onResults?.([...completed]); } diff --git a/ui/lib/provider-certification.ts b/ui/lib/provider-certification.ts index 8cdcebc1..1b9a0d20 100644 --- a/ui/lib/provider-certification.ts +++ b/ui/lib/provider-certification.ts @@ -1,3 +1,4 @@ +import { getApiErrorMessage } from '@/lib/api/errors'; import type { CertificationPath, CertificationStatus, @@ -113,4 +114,4 @@ export const getCertificationResultStatus = ( }; export const getErrorMessage = (error: unknown): string => - error instanceof Error ? error.message : 'Certification request failed'; + getApiErrorMessage(error, 'Certification request failed');