From c4e708bb32e0d8f20a6be3e76ef2edbb4747dba7 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 3 Oct 2026 20:00:26 +0200 Subject: [PATCH] fix: harden provider certification and native pricing --- routstr/algorithm.py | 10 +- routstr/core/admin.py | 19 +- routstr/payment/models.py | 26 +- routstr/upstream/base.py | 17 +- routstr/upstream/certification.py | 208 +++++---- routstr/upstream/certification_cache.py | 259 +++++++---- routstr/upstream/certification_probe.py | 156 +++++++ routstr/upstream/helpers.py | 5 +- routstr/upstream/model_paths.py | 22 +- routstr/upstream/ppqai.py | 81 ++-- .../test_admin_upstream_provider_report.py | 2 +- tests/integration/test_certify_endpoint.py | 12 +- tests/unit/test_algorithm.py | 10 +- tests/unit/test_certification_cache.py | 26 +- tests/unit/test_certification_deep_review.py | 419 ++++++++++++++++++ tests/unit/test_certification_token_limit.py | 5 +- .../unit/test_native_cache_pricing_policy.py | 131 ++++++ tests/unit/test_ppq_pricing_provenance.py | 138 ++++++ 18 files changed, 1285 insertions(+), 261 deletions(-) create mode 100644 routstr/upstream/certification_probe.py create mode 100644 tests/unit/test_certification_deep_review.py create mode 100644 tests/unit/test_native_cache_pricing_policy.py create mode 100644 tests/unit/test_ppq_pricing_provenance.py diff --git a/routstr/algorithm.py b/routstr/algorithm.py index 6a896665..c29f2f4f 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -239,7 +239,10 @@ def create_model_mappings( override_row, provider_fee = overrides_by_key[model_key] try: model_to_use = _row_to_model( - override_row, apply_provider_fee=True, provider_fee=provider_fee + override_row, + apply_provider_fee=True, + provider_fee=provider_fee, + provider_type=upstream.provider_type, ) except Exception as exc: # Stored pricing is JSON from whatever wrote the row, so @@ -315,7 +318,10 @@ def create_model_mappings( try: model_to_use = _row_to_model( - override_row, apply_provider_fee=True, provider_fee=provider_fee + override_row, + apply_provider_fee=True, + provider_fee=provider_fee, + provider_type=upstream_for_override.provider_type, ) except Exception as exc: logger.warning( diff --git a/routstr/core/admin.py b/routstr/core/admin.py index a36aa7bf..19ef78bb 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -617,7 +617,10 @@ async def upsert_provider_model( await refresh_model_maps() await _refresh_provider_model_paths(provider_pk) return _row_to_model( - row, apply_provider_fee=True, provider_fee=provider.provider_fee + row, + apply_provider_fee=True, + provider_fee=provider.provider_fee, + provider_type=provider.provider_type, ).dict() # type: ignore @@ -653,7 +656,10 @@ async def get_provider_model(provider_id: str, model_id: str) -> dict[str, objec # is not a usable number must be shown as it is, not encoded as `null`. return json_compliant( # type: ignore[return-value] _row_to_model( - row, apply_provider_fee=False, provider_fee=provider.provider_fee + row, + apply_provider_fee=False, + provider_fee=provider.provider_fee, + provider_type=provider.provider_type, ).dict() ) @@ -1311,7 +1317,10 @@ def _evaluate_model_row( ) -> _ModelEvaluation: try: configured: Model | None = _build_model_from_row( - row, apply_provider_fee=True, provider_fee=provider.provider_fee + row, + apply_provider_fee=True, + provider_fee=provider.provider_fee, + provider_type=provider.provider_type, ) build_error = None except Exception as exc: @@ -1352,7 +1361,9 @@ def _aggregate_row( """ evidence: dict[str, object] = {"checked": checked, "flagged": list(flagged)} if checked == 0: - return _report_row(row_id, "ok", title, empty_detail, evidence) + return _report_row( + row_id, "warn", title, f"Not evaluated — {empty_detail}", evidence + ) if flagged: return _report_row(row_id, fail_status, title, flagged_detail, evidence) return _report_row(row_id, "ok", title, ok_detail, evidence) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 3f17eb7a..f358d2e4 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -316,8 +316,16 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis return [] +def allows_cache_pricing_backfill(provider_type: str | None) -> bool: + return provider_type not in {"ppqai", "venice"} + + def _build_model_from_row( - row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 + row: ModelRow, + apply_provider_fee: bool = False, + provider_fee: float = 1.01, + *, + provider_type: str | None = None, ) -> Model: """The deterministic USD view of a stored model row, before the sats conversion.""" architecture = json.loads(row.architecture) @@ -346,7 +354,8 @@ def _build_model_from_row( # forwarded_model_id="deepseek-v4-flash") would otherwise look up the alias # and miss the cache rate. pricing_model_id = getattr(row, "forwarded_model_id", None) or row.id - parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing) + if allows_cache_pricing_backfill(provider_type): + parsed_pricing = backfill_cache_pricing(pricing_model_id, parsed_pricing) if apply_provider_fee: parsed_pricing = Pricing.parse_obj( @@ -383,9 +392,15 @@ def _build_model_from_row( def _row_to_model( - row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01 + row: ModelRow, + apply_provider_fee: bool = False, + provider_fee: float = 1.01, + *, + provider_type: str | None = None, ) -> Model: - model = _build_model_from_row(row, apply_provider_fee, provider_fee) + model = _build_model_from_row( + row, apply_provider_fee, provider_fee, provider_type=provider_type + ) try: sats_to_usd = sats_usd_price() @@ -430,6 +445,9 @@ async def list_models( provider_fee=providers_by_id[r.upstream_provider_id].provider_fee if r.upstream_provider_id in providers_by_id else 1.01, + provider_type=providers_by_id[r.upstream_provider_id].provider_type + if r.upstream_provider_id in providers_by_id + else None, ) except Exception as e: # Stored pricing/architecture is JSON from whatever wrote the row, so diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 4c4b68e8..1b1310e1 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -5699,12 +5699,17 @@ class BaseUpstreamProvider: code=error_code, ) + @property + def allow_cache_pricing_backfill(self) -> bool: + from ..payment.models import allows_cache_pricing_backfill + + return allows_cache_pricing_backfill(self.provider_type) + def _apply_provider_fee_to_model(self, model: Model) -> Model: """Apply provider fee to model's USD pricing and calculate max costs. - Cache rates missing from the upstream pricing feed are backfilled from - litellm's cost map first, so they carry the provider fee like every - other price component. + Providers with native price catalogs can disable generic cache-rate + backfill to avoid treating another provider's rates as their own. Args: model: Model object to update @@ -5712,7 +5717,11 @@ class BaseUpstreamProvider: Returns: Model with provider fee applied to pricing and max costs calculated """ - base_pricing = backfill_cache_pricing(model.id, model.pricing) + base_pricing = ( + backfill_cache_pricing(model.id, model.pricing) + if self.allow_cache_pricing_backfill + else model.pricing + ) adjusted_pricing = Pricing.parse_obj( {k: v * self.provider_fee for k, v in base_pricing.dict().items()} ) diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 9b00660b..a89b5a93 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -1,12 +1,12 @@ """Live certification checks for an upstream provider endpoint. Extends the read-only pricing rows, which never touch the network, with the -ones that must: a ``/models`` heartbeat and a one-token completion. +ones that must: a ``/models`` heartbeat and a bounded completion. Probes call the upstream directly with ``httpx``, never through the node's -billing path — no reservation, no Cashu. Upstream spend is one one-token -completion, plus two or three one-token completions on a ~4.4k-token prompt -when the cache checks are enabled (per certified model path). +billing path — no reservation, no Cashu. Output budgets start at 32 tokens +and rise only on recognized limit rejections, up to 2048. Cache checks add +repeated completions on a ~4.4k-token prompt (per certified model path). They sit behind ``POST …/certify`` rather than the read-only ``GET …/report`` because they can block for the length of the timeout. """ @@ -21,7 +21,7 @@ import os import sys import time from collections.abc import Callable -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any from urllib.parse import urlparse @@ -31,6 +31,14 @@ from ..core.logging import get_logger from ..payment.cost_calculation import _resolve_usd_cost, calculate_cost from ..payment.rates import coerce_rate from ..payment.usage import normalize_usage +from .certification_probe import ( + PROBE_MAX_TOKENS, + PROBE_TIMEOUT_SECONDS, + send_completion, +) +from .certification_probe import ( + wants_max_completion_tokens as wants_max_completion_tokens, +) from .model_paths import is_openrouter_base_url if TYPE_CHECKING: @@ -45,12 +53,7 @@ STATUS_FAIL = "fail" TICKS = {STATUS_OK: "☑️", STATUS_WARN: "⚠️", STATUS_FAIL: "❌"} -# Bounded so a dead upstream fails the row rather than wedging the request. -PROBE_TIMEOUT_SECONDS = 15.0 - -# The cheapest request that still exercises the usage/cost path. -PROBE_MAX_TOKENS = 1 -PROBE_PROMPT = "ping" +PROBE_PROMPT = "Reply with the single word: ok" # ``calculate_cost`` demands a reservation ceiling; any value at or above the # real charge behaves identically. @@ -128,12 +131,12 @@ CHECKLIST_GOALS: tuple[tuple[str, str, tuple[str, ...]], ...] = ( ), ( "caching", - "Prompt caching — cache hits reported and billed at the cache rate", + "Prompt caching — hits reported and cached cost calculation verified", ("cache.reported", "cache.billing"), ), ( "margin", - "Margin — node charge covers the upstream's cost", + "Pricing target — token estimate covers the fee-adjusted target", ("cost.margin",), ), ) @@ -166,7 +169,7 @@ def build_checklist(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: @dataclass class ProbeResult: - """Raw outcome of the two live HTTP calls a probe makes.""" + """Outcome of model discovery and the bounded completion attempts.""" base_url: str models_url: str @@ -180,19 +183,11 @@ class ProbeResult: chat_payload: dict[str, Any] | None = None chat_error: str | None = None chat_latency_ms: float | None = None - # ``max_completion_tokens`` once the upstream rejected ``max_tokens``. token_limit_field: str = "max_tokens" - - -def wants_max_completion_tokens(status: int | None, payload: Any) -> bool: - """Whether a 400 names ``max_completion_tokens`` as the field to use. - - OpenAI's o-series and gpt-5 reject ``max_tokens`` on chat completions with - "Unsupported parameter: 'max_tokens' ... Use 'max_completion_tokens'". - """ - if status != 400 or payload is None: - return False - return "max_completion_tokens" in json.dumps(payload, default=str) + token_limit: int = PROBE_MAX_TOKENS + chat_attempts: list[dict[str, Any]] = field(default_factory=list) + provider_type: str | None = None + completion_skip_reason: str | None = None @dataclass @@ -269,12 +264,11 @@ async def probe_upstream( upstream: "BaseUpstreamProvider | None" = None, model: "Model | None" = None, ) -> ProbeResult: - """Call the upstream's ``/models`` and a one-token completion. + """Discover models and probe supported chat APIs with bounded correction. - A completion refused with a 400 naming ``max_completion_tokens`` is - retried once with that field (OpenAI o-series, gpt-5). Each HTTP call, - including its body read, has an elapsed-time deadline. A transport - failure is a ``fail`` row, not a failed admin request. + Each HTTP call, including its body read, has an elapsed-time deadline. + Transport failures are rows, not failed admin requests. Native System One + inference is deliberately unverified until a supported fixture exists. """ shape = probe_shape(base_url, api_key, upstream, model) result = ProbeResult( @@ -282,7 +276,14 @@ async def probe_upstream( models_url=shape.models_url, chat_url=shape.chat_url, endpoint_tag=endpoint_tag, + provider_type=upstream.provider_type if upstream is not None else None, ) + if result.provider_type == "typesafe": + result.completion_skip_reason = ( + "Not evaluated — TypeSafe uses native /v1/systemone, not chat " + "completions. No supported native inference fixture is configured; " + "usage, billing and cache behavior remain unverified." + ) headers = shape.headers owns_client = client is None @@ -313,22 +314,11 @@ async def probe_upstream( result.models_error = f"{type(exc).__name__}: {exc}" result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) - if not model_id: + if not model_id or result.completion_skip_reason: return result await _probe_chat( client, result, model_id, shape, timeout, upstream, model, "max_tokens" ) - if wants_max_completion_tokens(result.chat_status, result.chat_payload): - await _probe_chat( - client, - result, - model_id, - shape, - timeout, - upstream, - model, - "max_completion_tokens", - ) finally: if owns_client: await client.aclose() @@ -346,7 +336,7 @@ async def _probe_chat( model: "Model | None", token_field: str, ) -> None: - """Send the one-token completion and record its outcome on ``result``.""" + """Record the effective provider-shaped field, budget and every attempt.""" request_body: dict[str, Any] = { "model": model_id, "messages": [{"role": "user", "content": PROBE_PROMPT}], @@ -358,35 +348,21 @@ async def _probe_chat( "order": [result.endpoint_tag], "allow_fallbacks": False, } - result.token_limit_field = token_field - result.chat_status = None - result.chat_payload = None - result.chat_error = None - started = time.monotonic() - try: - async with asyncio.timeout(timeout): - response = await client.post( - result.chat_url, - json=shape_body(request_body, upstream, model), - headers=shape.headers, - params=shape.chat_params, - ) - result.chat_status = response.status_code - result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) - try: - payload = response.json() - except Exception as exc: # noqa: BLE001 - any decode failure is the signal - result.chat_error = f"{type(exc).__name__}: {exc}" - else: - if isinstance(payload, dict): - result.chat_payload = payload - else: - result.chat_error = ( - f"expected a JSON object, got {type(payload).__name__}" - ) - except Exception as exc: # noqa: BLE001 - transport failure is a row status - result.chat_error = f"{type(exc).__name__}: {exc}" - result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2) + outcome = await send_completion( + client, + result.chat_url, + shape_body(request_body, upstream, model), + shape.headers, + shape.chat_params, + timeout, + ) + result.token_limit_field = outcome.token_limit_field + result.token_limit = outcome.token_limit + result.chat_status = outcome.status + result.chat_payload = outcome.payload + result.chat_error = outcome.error + result.chat_latency_ms = outcome.latency_ms + result.chat_attempts = outcome.attempts # Row builders are pure: the network lives only in ``probe_upstream`` and @@ -473,13 +449,16 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]: {"url": probe.models_url, "error": probe.models_error}, ) - data = payload.get("data") + data_key, id_key = ( + ("models", "name") if probe.provider_type == "typesafe" else ("data", "id") + ) + data = payload.get(data_key) if not isinstance(data, list): return certification_row( "endpoint.models_payload", STATUS_FAIL, "Models payload has the expected shape", - f'Expected a top-level "data" list, got {type(data).__name__}.', + f'Expected a top-level "{data_key}" list, got {type(data).__name__}.', { "url": probe.models_url, "top_level_keys": sorted(payload.keys()), @@ -487,9 +466,9 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]: ) ids = [ - item["id"] + item[id_key] for item in data - if isinstance(item, dict) and isinstance(item.get("id"), str) and item["id"] + if isinstance(item, dict) and isinstance(item.get(id_key), str) and item[id_key] ] evidence: dict[str, Any] = { "url": probe.models_url, @@ -502,15 +481,15 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]: "endpoint.models_payload", STATUS_FAIL, "Models payload has the expected shape", - f'The "data" list carries no entry with a non-empty string "id" ' - f"({len(data)} entries).", + f'The "{data_key}" list carries no entry with a non-empty string ' + f'"{id_key}" ({len(data)} entries).', evidence, ) return certification_row( "endpoint.models_payload", STATUS_OK, "Models payload has the expected shape", - f"{len(ids)} of {len(data)} entries carry a string id.", + f"{len(ids)} of {len(data)} entries carry a string {id_key}.", evidence, ) @@ -518,14 +497,25 @@ def models_payload_row(probe: ProbeResult) -> dict[str, Any]: def usage_capture_row(probe: ProbeResult) -> dict[str, Any]: """Check a completion comes back with token usage the node can bill on. - A missing ``usage`` object means the node has nothing to price and the - request settles for free. Broken, but still usable, so ``warn``. + Missing upstream usage is a coverage gap. The proxy can estimate usage + before settlement; this direct probe does not exercise that fallback. """ evidence: dict[str, Any] = { "url": probe.chat_url, "status_code": probe.chat_status, "latency_ms": probe.chat_latency_ms, + "token_limit_field": probe.token_limit_field, + "token_limit": probe.token_limit, + "attempts": probe.chat_attempts, } + if probe.completion_skip_reason: + return certification_row( + "usage.capture", + STATUS_WARN, + "Token usage captured from a completion", + probe.completion_skip_reason, + evidence, + ) if probe.chat_status is None: evidence["error"] = probe.chat_error return certification_row( @@ -536,13 +526,17 @@ def usage_capture_row(probe: ProbeResult) -> dict[str, Any]: evidence, ) if not 200 <= probe.chat_status < 300: - evidence["body"] = _truncate(probe.chat_payload) + evidence["body"] = ( + probe.chat_attempts[-1]["body"] + if probe.chat_attempts + else _truncate(probe.chat_payload) + ) return certification_row( "usage.capture", STATUS_FAIL, "Token usage captured from a completion", f"{probe.chat_url} answered {probe.chat_status} for a " - f"{PROBE_MAX_TOKENS}-token probe.", + f"{probe.token_limit}-token probe ({probe.token_limit_field}).", evidence, ) if probe.chat_payload is None or not isinstance(probe.chat_payload, dict): @@ -576,8 +570,9 @@ def usage_capture_row(probe: ProbeResult) -> dict[str, Any]: "usage.capture", STATUS_WARN, "Token usage captured from a completion", - 'The completion carried no "usage" object, so the node has no ' - "token counts to bill on and the request would settle as (0+0).", + 'The completion carried no "usage" object. The proxy may estimate ' + "missing usage before settlement; this probe does not verify that " + "fallback or the resulting client debit.", evidence, ) evidence["input_tokens"] = normalized.input_tokens @@ -782,8 +777,8 @@ def cost_prompt_completion_row( "cost.prompt_completion", STATUS_FAIL, "Prompt and completion cost calculated", - f"The expected charge could not be derived from the configured " - f"pricing: {type(exc).__name__}: {exc}.", + f"The expected charge could not be derived from the selected " + f"billing basis: {type(exc).__name__}: {exc}.", evidence, ) @@ -826,12 +821,17 @@ def cost_prompt_completion_row( ): mismatches.append(f"input {actual_input} != {expected_input}") + basis_label = ( + "upstream-reported USD converted with the provider fee and exchange rate" + if basis == "upstream_reported_usd" + else "configured token pricing" + ) if mismatches: return certification_row( "cost.prompt_completion", STATUS_FAIL, "Prompt and completion cost calculated", - "The computed charge disagrees with the configured pricing: " + f"The computed charge disagrees with {basis_label}: " + "; ".join(mismatches) + ".", evidence, @@ -840,10 +840,10 @@ def cost_prompt_completion_row( "cost.prompt_completion", STATUS_OK, "Prompt and completion cost calculated", - f"Charged {actual_total} msats ({actual_input} input + " + f"Calculated {actual_total} msats ({actual_input} input + " f"{actual_output} output) for {usage.input_tokens} prompt and " - f"{usage.output_tokens} completion tokens, matching the configured " - f"pricing.", + f"{usage.output_tokens} completion tokens, matching {basis_label}. " + "This is a calculation, not a measured client debit.", evidence, ) @@ -901,8 +901,27 @@ async def run_live_checks( ), ] + from .certification_cache import run_cache_checks, skipped_cache_rows + + if probe.completion_skip_reason: + rows.append( + certification_row( + "cost.prompt_completion", + STATUS_WARN, + "Prompt and completion cost calculated", + probe.completion_skip_reason, + {}, + ) + ) + rows.extend(skipped_cache_rows(probe.completion_skip_reason)) + return rows + cost_data: Any = None - if probe.chat_payload is not None and probe.chat_status is not None: + if ( + probe.chat_payload is not None + and probe.chat_status is not None + and 200 <= probe.chat_status < 300 + ): try: cost_data = await calculate_cost( probe.chat_payload, @@ -939,8 +958,6 @@ async def run_live_checks( ) ) - from .certification_cache import run_cache_checks, skipped_cache_rows - if not check_cache: rows.extend(skipped_cache_rows("Skipped — cache checks disabled.")) elif ( @@ -973,6 +990,7 @@ async def run_live_checks( endpoint_tag=endpoint_tag, upstream=upstream, token_limit_field=probe.token_limit_field, + token_limit=probe.token_limit, ) ) return rows diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index 24751712..2b8eeba5 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -1,10 +1,10 @@ """Prompt-cache and margin certification for an upstream provider. -Three questions the one-token probe cannot answer: +Three questions the short completion probe cannot answer: * does the upstream *report* prompt-cache hits in a dialect the node parses, -* does the node bill cached reads at the discounted rate (client side), and -* does the node's charge cover what the upstream charged (node side). +* does the cost engine calculate cached-completion charges correctly, and +* does the token estimate cover the reported upstream cost target. The cache probe sends the same long system prompt twice; the second call is the one expected to report cached reads. Calls go straight to the upstream @@ -13,8 +13,6 @@ 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 @@ -36,11 +34,13 @@ from .certification import ( _fixed_token_pricing_active, _reported_usd_cost, _token_rates, + _truncate, certification_row, probe_shape, safe_row, shape_body, ) +from .certification_probe import CompletionProbe, send_completion if TYPE_CHECKING: from ..payment.models import Model @@ -58,8 +58,8 @@ ROW_BILLING = "cache.billing" ROW_MARGIN = "cost.margin" TITLE_REPORTED = "Upstream reports prompt-cache hits" -TITLE_BILLING = "Cached tokens billed at the cache-read rate" -TITLE_MARGIN = "Node charge covers upstream cost" +TITLE_BILLING = "Cached completion cost calculated" +TITLE_MARGIN = "Token estimate covers fee-adjusted target" def cache_probe_prefix() -> str: @@ -86,6 +86,9 @@ class CacheProbeResult: payloads: list[dict[str, Any] | None] = field(default_factory=list) errors: list[str | None] = field(default_factory=list) latencies_ms: list[float | None] = field(default_factory=list) + attempts: list[dict[str, Any]] = field(default_factory=list) + token_limit_field: str = "max_tokens" + token_limit: int = PROBE_MAX_TOKENS @property def second_payload(self) -> dict[str, Any] | None: @@ -104,6 +107,7 @@ def _request_body( fmt: str, endpoint_tag: str | None, token_field: str = "max_tokens", + token_limit: int = PROBE_MAX_TOKENS, ) -> dict[str, Any]: system: Any if fmt == "cache_control": @@ -122,7 +126,7 @@ def _request_body( {"role": "system", "content": system}, {"role": "user", "content": CACHE_PROBE_QUESTION}, ], - token_field: PROBE_MAX_TOKENS, + token_field: token_limit, "stream": False, } if endpoint_tag: @@ -133,45 +137,11 @@ def _request_body( return body -async def _post_completion( - client: httpx.AsyncClient, - url: str, - body: Any, - headers: dict[str, str], - params: dict[str, str], - timeout: float, -) -> tuple[int | None, dict[str, Any] | None, str | None, float]: - started = time.monotonic() - try: - async with asyncio.timeout(timeout): - response = await client.post(url, json=body, headers=headers, params=params) - except Exception as exc: # noqa: BLE001 - transport failure is a row status - latency = round((time.monotonic() - started) * 1000, 2) - return None, None, f"{type(exc).__name__}: {exc}", latency - latency = round((time.monotonic() - started) * 1000, 2) - try: - payload = response.json() - except Exception as exc: # noqa: BLE001 - any decode failure is the signal - return response.status_code, None, f"{type(exc).__name__}: {exc}", latency - if not isinstance(payload, dict): - return ( - response.status_code, - None, - f"expected a JSON object, got {type(payload).__name__}", - latency, - ) - return response.status_code, payload, None, latency - - -def _record( - result: CacheProbeResult, - outcome: tuple[int | None, dict[str, Any] | None, str | None, float], -) -> None: - status, payload, error, latency = outcome - result.statuses.append(status) - result.payloads.append(payload) - result.errors.append(error) - result.latencies_ms.append(latency) +def _record(result: CacheProbeResult, outcome: CompletionProbe) -> None: + result.statuses.append(outcome.status) + result.payloads.append(outcome.payload) + result.errors.append(outcome.error) + result.latencies_ms.append(outcome.latency_ms) def _is_2xx(status: int | None) -> bool: @@ -189,6 +159,7 @@ async def probe_cache( upstream: "BaseUpstreamProvider | None" = None, model: "Model | None" = None, token_limit_field: str = "max_tokens", + token_limit: int = PROBE_MAX_TOKENS, ) -> CacheProbeResult: """Send the same long prompt twice. @@ -198,17 +169,27 @@ async def probe_cache( succeeded. Each call's elapsed deadline includes the response body. """ shape = probe_shape(base_url, api_key, upstream, model) - result = CacheProbeResult(chat_url=shape.chat_url, endpoint_tag=endpoint_tag) + result = CacheProbeResult( + chat_url=shape.chat_url, + endpoint_tag=endpoint_tag, + token_limit_field=token_limit_field, + token_limit=token_limit, + ) prefix = cache_probe_prefix() owns_client = client is None http = client if client is not None else httpx.AsyncClient(timeout=timeout) - async def post( - fmt: str, - ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: - body = _request_body(model_id, prefix, fmt, endpoint_tag, token_limit_field) - return await _post_completion( + async def post(fmt: str, phase: str) -> CompletionProbe: + body = _request_body( + model_id, + prefix, + fmt, + endpoint_tag, + result.token_limit_field, + result.token_limit, + ) + outcome = await send_completion( http, result.chat_url, shape_body(body, upstream, model), @@ -216,16 +197,23 @@ async def probe_cache( shape.chat_params, timeout, ) + result.token_limit_field = outcome.token_limit_field + result.token_limit = outcome.token_limit + result.attempts.extend( + {**attempt, "request_format": fmt, "phase": phase} + for attempt in outcome.attempts + ) + return outcome try: - first = await post("cache_control") - if first[0] in (400, 422): + first = await post("cache_control", "warm_up") + if first.status in (400, 422): result.request_format = "plain" - first = await post("plain") + first = await post("plain", "warm_up") _record(result, first) - if not _is_2xx(first[0]): + if not _is_2xx(first.status): return result - second = await post(result.request_format) + second = await post(result.request_format, "repeated") _record(result, second) finally: if owns_client: @@ -262,16 +250,29 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]: "endpoint_tag": probe.endpoint_tag, "statuses": probe.statuses, "latencies_ms": probe.latencies_ms, + "attempts": probe.attempts, + "errors": probe.errors, + "token_limit_field": probe.token_limit_field, + "token_limit": probe.token_limit, } payload = probe.second_payload if payload is None or not _is_2xx(probe.statuses[-1] if probe.statuses else None): - evidence["error"] = probe.second_error + status = probe.statuses[-1] if probe.statuses else None + if status is not None and not _is_2xx(status): + body = ( + probe.attempts[-1].get("body") + if probe.attempts + else _truncate(probe.payloads[-1] if probe.payloads else None) + ) + reason = f"HTTP {status}" + (f": {body}" if body else "") + else: + reason = probe.second_error or "no decodable completion body" + evidence["error"] = reason return certification_row( ROW_REPORTED, STATUS_FAIL, TITLE_REPORTED, - f"The cache probe did not get two successful completions: " - f"{probe.second_error or 'no response body'}.", + f"The cache probe did not get two successful completions: {reason}.", evidence, ) @@ -308,17 +309,17 @@ def cache_reported_row(probe: CacheProbeResult) -> dict[str, Any]: STATUS_FAIL, TITLE_REPORTED, "The upstream reported cache tokens under fields the node does not " - f"parse ({', '.join(raw_keys)}); cached reads would be billed at " - "the full input rate.", + f"parse ({', '.join(raw_keys)}). Token-based pricing cannot apply " + "a cache-read rate to those fields; reported-USD billing may take " + "precedence.", evidence, ) return certification_row( ROW_REPORTED, STATUS_WARN, TITLE_REPORTED, - "Two identical prompts produced no cache hit. Either the model does " - "not support prompt caching or the upstream hides it; clients pay the " - "full input rate on repeated prompts.", + "Two identical prompts produced no reported cache hit. Caching may " + "be unsupported or unreported; no cache discount is verified.", evidence, ) @@ -329,6 +330,8 @@ def cache_billing_row( probe: CacheProbeResult, cost_data: Any, pricing_known: bool = True, + provider_fee: float = 1.0, + sats_to_usd: float | None = None, ) -> dict[str, Any]: usage = _usage_of(probe.second_payload) evidence: dict[str, Any] = {"model_id": model.id} @@ -345,7 +348,7 @@ def cache_billing_row( ROW_BILLING, STATUS_WARN, TITLE_BILLING, - "No pricing is known for this model, so the cache discount cannot " + "No pricing is known for this model, so the cache charge cannot " "be verified.", evidence, ) @@ -395,14 +398,59 @@ def cache_billing_row( } ) + evidence["basis"] = ( + "upstream_reported_usd" if reported_usd > 0 else "configured_token_pricing" + ) if reported_usd > 0: + evidence.update( + { + "token_estimate_msats": expected_total, + "provider_fee": provider_fee, + "sats_usd_price": sats_to_usd, + "expected_total_msats": None, + "upstream_discount_verified": False, + } + ) + if sats_to_usd is None: + return certification_row( + ROW_BILLING, + STATUS_WARN, + TITLE_BILLING, + "No exchange rate was supplied, so the reported-USD calculation " + "cannot be verified.", + evidence, + ) + try: + expected_total = _expected_usd_msats( + reported_usd, provider_fee, sats_to_usd + ) + except (ValueError, OverflowError) as exc: + evidence["error"] = f"{type(exc).__name__}: {exc}" + return certification_row( + ROW_BILLING, + STATUS_FAIL, + TITLE_BILLING, + f"The expected reported-USD charge could not be derived: {exc}.", + evidence, + ) + evidence["expected_total_msats"] = expected_total + if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS: + return certification_row( + ROW_BILLING, + STATUS_FAIL, + TITLE_BILLING, + f"The engine calculated {actual_total} msats but reported USD, " + f"the provider fee and exchange rate imply {expected_total} msats.", + evidence, + ) return certification_row( ROW_BILLING, STATUS_OK, TITLE_BILLING, - f"Billed {actual_total} msats from the upstream-reported cost, " - f"which already carries the cache discount " - f"(full token price would be {full_total} msats).", + f"Calculated {actual_total} msats, matching the independently " + "converted upstream-reported USD with the provider fee. " + f"The configured full-input comparator is {full_total} msats; " + "an upstream cache discount and actual client debit are not verified.", evidence, ) if abs(actual_total - expected_total) > COST_TOLERANCE_MSATS: @@ -410,7 +458,7 @@ def cache_billing_row( ROW_BILLING, STATUS_FAIL, TITLE_BILLING, - f"The engine charged {actual_total} msats but the configured cache " + f"The engine calculated {actual_total} msats but the configured cache " f"rate implies {expected_total} msats.", evidence, ) @@ -424,16 +472,16 @@ def cache_billing_row( ROW_BILLING, STATUS_WARN, TITLE_BILLING, - f"Cached reads are billed at the full input rate ({actual_total} " - f"msats) because {reason}; clients pay more than the upstream " - "charges.", + f"Cached reads use the full input rate ({actual_total} calculated " + f"msats) because {reason}. Upstream cost, any upstream discount " + "and actual client debit are unverified.", evidence, ) return certification_row( ROW_BILLING, STATUS_OK, TITLE_BILLING, - f"Charged {actual_total} msats for {usage.cache_read_tokens} cached " + f"Calculated {actual_total} msats for {usage.cache_read_tokens} cached " f"tokens, {full_total - actual_total} msats below the full input price.", evidence, ) @@ -449,12 +497,9 @@ def cost_margin_row( ) -> dict[str, Any]: """Configured token pricing must cover what the upstream reports charging. - Responses that carry a USD cost are billed from it, so they cannot lose - money themselves; they are used here as a price sample. The configured - token pricing is what every other path bills from (streams, upstreams - that omit cost, the served ``/v1/models`` list), so a sample where it - falls below the fee-adjusted upstream cost means those paths underprice. - Upstreams that report no cost give no sample and the row stays a warn. + Reported USD takes precedence in the cost engine; the token-price estimate + is a separate fallback comparison, not a measured debit. A fee-target gap + need not be a loss against raw upstream cost. Missing cost stays unverified. ``model`` carries the pricing the proxy reserves and token-bills with, which on a pinned endpoint is that endpoint's own rates. @@ -477,6 +522,7 @@ def cost_margin_row( samples: list[dict[str, Any]] = [] short: list[str] = [] + below_raw: list[str] = [] for payload in payloads: if not isinstance(payload, dict): continue @@ -489,6 +535,7 @@ def cost_margin_row( upstream_total = _expected_usd_msats( reported_usd, provider_fee, sats_to_usd ) + raw_upstream_total = _expected_usd_msats(reported_usd, 1.0, sats_to_usd) except (ValueError, OverflowError) as exc: evidence["error"] = f"{type(exc).__name__}: {exc}" return certification_row( @@ -502,11 +549,17 @@ def cost_margin_row( "usage": usage.dict(), "reported_usd": reported_usd, "upstream_msats_with_fee": upstream_total, + "upstream_msats": raw_upstream_total, "configured_msats": configured_total, + "token_estimate_minus_upstream_msats": configured_total + - raw_upstream_total, + "token_estimate_minus_fee_target_msats": configured_total - upstream_total, } samples.append(sample) if configured_total + COST_TOLERANCE_MSATS < upstream_total: short.append(f"{configured_total} < {upstream_total}") + if configured_total + COST_TOLERANCE_MSATS < raw_upstream_total: + below_raw.append(f"{configured_total} < {raw_upstream_total}") evidence["samples"] = samples if not samples: @@ -524,17 +577,30 @@ def cost_margin_row( ROW_MARGIN, STATUS_FAIL, TITLE_MARGIN, - "Configured pricing is below the upstream's reported cost " - f"(configured < upstream msats: {'; '.join(short)}); token-billed " - "requests lose money.", + "The token-pricing estimate is below the fee-adjusted cost target " + f"(estimate < target msats: {'; '.join(short)}). " + + ( + "It is also below raw upstream cost " + f"(estimate < raw msats: {'; '.join(below_raw)}). " + if below_raw + else "Raw upstream cost is covered; the gap is in the fee target. " + ) + + "Responses reporting USD use reported-cost billing instead; " + "this check does not measure client debits.", evidence, ) return certification_row( ROW_MARGIN, STATUS_OK, TITLE_MARGIN, - f"Configured pricing covers the upstream's reported cost on " - f"{len(samples)} sampled completion(s).", + f"Token pricing covers the fee-adjusted cost target on " + f"{len(samples)} sampled completion(s)." + + ( + " The fee multiplier is below break-even: raw upstream cost exceeds " + f"the token estimate on samples ({'; '.join(below_raw)})." + if below_raw + else "" + ), evidence, ) @@ -570,6 +636,7 @@ async def run_cache_checks( endpoint_tag: str | None = None, upstream: "BaseUpstreamProvider | None" = None, token_limit_field: str = "max_tokens", + token_limit: int = PROBE_MAX_TOKENS, ) -> list[dict[str, Any]]: """Run the cache probe and build the three cache/margin rows.""" probe = await probe_cache( @@ -582,8 +649,15 @@ async def run_cache_checks( upstream=upstream, model=model, token_limit_field=token_limit_field, + token_limit=token_limit, + ) + cost_data = await _price_payload( + probe.second_payload + if probe.statuses and _is_2xx(probe.statuses[-1]) + else None, + model, + provider_fee, ) - cost_data = await _price_payload(probe.second_payload, model, provider_fee) return [ safe_row(ROW_REPORTED, TITLE_REPORTED, lambda: cache_reported_row(probe)), safe_row( @@ -594,6 +668,8 @@ async def run_cache_checks( probe=probe, cost_data=cost_data, pricing_known=pricing_known, + provider_fee=provider_fee, + sats_to_usd=sats_to_usd, ), ), safe_row( @@ -601,7 +677,14 @@ async def run_cache_checks( TITLE_MARGIN, lambda: cost_margin_row( model=model, - payloads=[probe_payload, *probe.payloads], + payloads=[ + probe_payload, + *[ + payload + for status, payload in zip(probe.statuses, probe.payloads) + if _is_2xx(status) + ], + ], provider_fee=provider_fee, sats_to_usd=sats_to_usd, pricing_known=pricing_known, diff --git a/routstr/upstream/certification_probe.py b/routstr/upstream/certification_probe.py new file mode 100644 index 00000000..8a1f4b50 --- /dev/null +++ b/routstr/upstream/certification_probe.py @@ -0,0 +1,156 @@ +"""Bounded completion probes shared by ordinary and cache certification.""" + +from __future__ import annotations + +import asyncio +import json +import re +import time +from dataclasses import dataclass, field +from typing import Any + +import httpx + +PROBE_TIMEOUT_SECONDS = 15.0 +PROBE_TOKEN_BUDGETS = (32, 128, 512, 2048) +PROBE_MAX_TOKENS = PROBE_TOKEN_BUDGETS[0] +# One field correction plus the four bounded budgets, never general retries. +PROBE_MAX_ATTEMPTS = len(PROBE_TOKEN_BUDGETS) + 1 + + +def wants_max_completion_tokens(status: int | None, payload: Any) -> bool: + if status != 400 or payload is None: + return False + text = json.dumps(payload, default=str).lower() + return "max_completion_tokens" in text and ( + "unsupported" in text + and "max_tokens" in text + or "use 'max_completion_tokens'" in text + ) + + +def next_probe_budget(status: int | None, payload: Any, current: int) -> int | None: + if status not in (400, 422) or payload is None: + return None + text = json.dumps(payload, default=str).lower() + minimum = re.search( + r"(?:max_tokens|max_completion_tokens|max_output_tokens)" + r".{0,80}?(?:at least|minimum(?: of)?|>=|greater than or equal to)" + r"\D{0,12}(\d+)", + text, + ) + required = int(minimum[1]) if minimum else None + if required is not None: + if required <= current: + return None + elif not ( + "max_tokens or model output limit was reached" in text + or ( + any(name in text for name in ("max_tokens", "max_completion_tokens")) + and any( + hint in text for hint in ("higher max", "exhausted", "increase max") + ) + ) + ): + return None + return next( + (n for n in PROBE_TOKEN_BUDGETS if n > current and n >= (required or 0)), + None, + ) + + +def diagnostic_text(value: Any, headers: dict[str, str], limit: int = 1000) -> str: + """Redact known request credentials before bounding diagnostic text.""" + text = value if isinstance(value, str) else json.dumps(value, default=str) + for name, secret in headers.items(): + if name.lower() in {"authorization", "api-key", "x-api-key"} and secret: + text = text.replace(secret, "[redacted]") + if name.lower() == "authorization" and " " in secret: + text = text.replace(secret.split(" ", 1)[1], "[redacted]") + return text if len(text) <= limit else text[:limit] + "…" + + +@dataclass +class CompletionProbe: + status: int | None = None + payload: dict[str, Any] | None = None + error: str | None = None + latency_ms: float = 0.0 + token_limit_field: str = "max_tokens" + token_limit: int = PROBE_MAX_TOKENS + attempts: list[dict[str, Any]] = field(default_factory=list) + + +async def send_completion( + client: httpx.AsyncClient, + url: str, + body: dict[str, Any], + headers: dict[str, str], + params: dict[str, str], + timeout: float, +) -> CompletionProbe: + """Adapt only recognized token-limit rejections, preserving every call. + + The body is already provider-shaped. This policy is exclusive to synthetic + certification requests; user request budgets are never changed here. + """ + body = dict(body) + result = CompletionProbe() + for _ in range(PROBE_MAX_ATTEMPTS): + token_field = ( + "max_completion_tokens" if "max_completion_tokens" in body else "max_tokens" + ) + budget = int(body[token_field]) + if not 1 <= budget <= PROBE_TOKEN_BUDGETS[-1]: + raise ValueError("certification token budget is outside the probe ceiling") + result.token_limit_field = token_field + result.token_limit = budget + result.status, result.payload, result.error = None, None, None + diagnostic: Any = None + started = time.monotonic() + try: + async with asyncio.timeout(timeout): + response = await client.post( + url, json=body, headers=headers, params=params + ) + result.status = response.status_code + diagnostic = response.text + try: + payload = response.json() + except Exception as exc: # noqa: BLE001 - decoding failure is evidence + result.error = f"{type(exc).__name__}: {exc}" + else: + if isinstance(payload, dict): + result.payload = payload + else: + result.error = ( + f"expected a JSON object, got {type(payload).__name__}" + ) + except Exception as exc: # noqa: BLE001 - cancellation still propagates + result.error = f"{type(exc).__name__}: {exc}" + result.latency_ms = round((time.monotonic() - started) * 1000, 2) + if result.error: + result.error = diagnostic_text(result.error, headers) + result.attempts.append( + { + "status_code": result.status, + "latency_ms": result.latency_ms, + "token_limit_field": token_field, + "token_limit": budget, + "error": result.error, + "body": diagnostic_text(diagnostic, headers) + if diagnostic is not None + and (result.error or not 200 <= (result.status or 0) < 300) + else None, + } + ) + if token_field == "max_tokens" and wants_max_completion_tokens( + result.status, result.payload + ): + body["max_completion_tokens"] = body.pop("max_tokens") + continue + next_budget = next_probe_budget(result.status, result.payload, budget) + if next_budget is None: + break + body[token_field] = next_budget + return result diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index c94733bc..be14a825 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -123,7 +123,10 @@ async def get_all_models_with_overrides( if model_key is not None and model_key in overrides_by_key: override_row, provider_fee = overrides_by_key[model_key] all_models[(model.id.lower(), provider_key)] = _row_to_model( - override_row, apply_provider_fee=True, provider_fee=provider_fee + override_row, + apply_provider_fee=True, + provider_fee=provider_fee, + provider_type=upstream.provider_type, ) elif model.enabled: all_models[(model.id.lower(), provider_key)] = model diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index f7e46d81..f639e734 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -824,7 +824,9 @@ async def refresh_model_paths_periodically( break -def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None: +def _price_in_sats( + model: dict[str, Any], provider_fee: float, provider_type: str | None = None +) -> None: """Run a path's USD rates through the ``/v1/models`` pricing pipeline. Metadata copied from the provider model cache is already priced. OpenRouter @@ -842,13 +844,16 @@ def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None: TopProvider, _calculate_usd_max_costs, _update_model_sats_pricing, + allows_cache_pricing_backfill, backfill_cache_pricing, ) from ..payment.price import sats_usd_price try: model_id = model.get("forwarded_model_id") or model["id"] - usd = backfill_cache_pricing(model_id, Pricing.parse_obj(pricing)) + usd = Pricing.parse_obj(pricing) + if allows_cache_pricing_backfill(provider_type): + usd = backfill_cache_pricing(model_id, usd) usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()}) priced = Model( id=model_id, @@ -921,10 +926,13 @@ def apply_model_path_pricing( metadata.get("pricing"), dict ): return model - pricing = backfill_cache_pricing( - model.forwarded_model_id or row.model_id, - Pricing.parse_obj(metadata["pricing"]), - ) + from ..payment.models import allows_cache_pricing_backfill + + pricing = Pricing.parse_obj(metadata["pricing"]) + if allows_cache_pricing_backfill(row.provider_type): + pricing = backfill_cache_pricing( + model.forwarded_model_id or row.model_id, pricing + ) pricing = Pricing.parse_obj( {key: float(value) * provider_fee for key, value in pricing.dict().items()} ) @@ -963,7 +971,7 @@ def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]: if not isinstance(model, dict): model = {} model.setdefault("id", row.model_id) - _price_in_sats(model, provider_fee) + _price_in_sats(model, provider_fee, row.provider_type) return { "path": row.path, "provider": { diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 65ccbb57..bfa2e8a2 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import math import random import time from dataclasses import dataclass, field @@ -226,6 +227,27 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): if ppqai_model.id in self.IGNORED_MODEL_IDS: continue + api_pricing = ppqai_model.pricing.api or {} + input_price = api_pricing.get("input_per_1M") + if input_price is None: + input_price = ppqai_model.pricing.input_per_1M_tokens + output_price = api_pricing.get("output_per_1M") + if output_price is None: + output_price = ppqai_model.pricing.output_per_1M_tokens + if any( + rate is None or not math.isfinite(rate) or rate < 0 + for rate in (input_price, output_price) + ): + logger.warning( + "Skipping PPQ model without valid native token prices", + extra={"model_id": ppqai_model.id}, + ) + continue + assert input_price is not None and output_price is not None + pricing = Pricing( + prompt=input_price / 1_000_000, + completion=output_price / 1_000_000, + ) or_model = next( ( model @@ -238,44 +260,20 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) if or_model: - input_price = None - if ppqai_model.pricing.api: - input_price = ppqai_model.pricing.api.get("input_per_1M") - elif ppqai_model.pricing.input_per_1M_tokens: - input_price = ppqai_model.pricing.input_per_1M_tokens - - if input_price is not None: - or_model.pricing.prompt = input_price / 1_000_000 - - output_price = None - if ppqai_model.pricing.api: - output_price = ppqai_model.pricing.api.get("output_per_1M") - elif ppqai_model.pricing.output_per_1M_tokens: - output_price = ppqai_model.pricing.output_per_1M_tokens - - if output_price is not None: - or_model.pricing.completion = output_price / 1_000_000 - - if cl := ppqai_model.context_length: - or_model.context_length = cl - models.append(or_model) + # OpenRouter supplies metadata, not PPQ billing rates. + models.append( + or_model.copy( + deep=True, + update={ + "id": ppqai_model.id, + "pricing": pricing, + "sats_pricing": None, + "context_length": ppqai_model.context_length + or or_model.context_length, + }, + ) + ) else: - input_price = 0.0 - if ppqai_model.pricing.api: - input_price = ppqai_model.pricing.api.get( - "input_per_1M", 0.0 - ) - elif ppqai_model.pricing.input_per_1M_tokens: - input_price = ppqai_model.pricing.input_per_1M_tokens - - output_price = 0.0 - if ppqai_model.pricing.api: - output_price = ppqai_model.pricing.api.get( - "output_per_1M", 0.0 - ) - elif ppqai_model.pricing.output_per_1M_tokens: - output_price = ppqai_model.pricing.output_per_1M_tokens - models.append( Model( id=ppqai_model.id, @@ -290,14 +288,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): tokenizer="Unknown", instruct_type=None, ), - pricing=Pricing( - prompt=input_price / 1_000_000, - completion=output_price / 1_000_000, - request=0.0, - image=0.0, - web_search=0.0, - internal_reasoning=0.0, - ), + pricing=pricing, ) ) except Exception as e: diff --git a/tests/integration/test_admin_upstream_provider_report.py b/tests/integration/test_admin_upstream_provider_report.py index c64fcaba..2b433bb6 100644 --- a/tests/integration/test_admin_upstream_provider_report.py +++ b/tests/integration/test_admin_upstream_provider_report.py @@ -448,7 +448,7 @@ async def test_cache_rate_ignores_an_enabled_model_that_is_not_served( assert resp.status_code == 200, resp.text row = _find_row(resp.json()["rows"], "pricing.cache_rate") - assert row["status"] == "ok", row + assert row["status"] == "warn", row assert row["evidence"]["checked"] == 0 assert "unserved-model" not in json.dumps(row["evidence"]) diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index 065f888f..c7eedf7b 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -343,9 +343,7 @@ async def test_certify_margin_bills_pinned_path_pricing( }, base_url=base_url, ) - model_path = encode_model_path( - base_url, "cert-test-model", "deepinfra/fp8" - ) + model_path = encode_model_path(base_url, "cert-test-model", "deepinfra/fp8") integration_session.add( ModelPathRow( model_id="cert-test-model", @@ -834,6 +832,14 @@ async def test_certify_explicit_discovered_model_without_override( assert _find_row(rows, "endpoint.reachable")["status"] == "ok" assert _find_row(rows, "usage.capture")["status"] == "ok" assert _find_row(rows, "cost.prompt_completion")["status"] == "ok" + pricing_rows = [row for row in rows if row["id"].startswith("pricing.")] + assert len(pricing_rows) == 4 + assert all(row["evidence"]["checked"] == 0 for row in pricing_rows) + assert all(row["status"] == "warn" for row in pricing_rows) + pricing_goal = next( + goal for goal in resp.json()["checklist"] if goal["goal"] == "pricing_v1_models" + ) + assert pricing_goal["status"] == "warn" @pytest.mark.integration diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index f5c875f9..42c52cf2 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -272,7 +272,10 @@ def test_create_model_mappings_applies_custom_provider_fees_before_advertising( 2: SimpleNamespace(id="shared-model", upstream_provider_id=2, enabled=True), } - def fake_row_to_model(row, *, apply_provider_fee, provider_fee) -> Model: # type: ignore[no-untyped-def] + def fake_row_to_model( + row, *, apply_provider_fee, provider_fee, provider_type + ) -> Model: # type: ignore[no-untyped-def] + assert provider_type == providers[row.upstream_provider_id - 1].provider_type assert apply_provider_fee is True return create_test_model( row.id, @@ -943,7 +946,10 @@ def test_create_model_mappings_uppercase_prefixed_base_keeps_top_tier() -> None: "prefixed", "https://prefixed.example/v1", db_id=1, models=[prefixed_cheap] ) forwarded_provider = create_test_provider( - "forwarded", "https://forwarded.example/v1", db_id=2, models=[forwarded_expensive] + "forwarded", + "https://forwarded.example/v1", + db_id=2, + models=[forwarded_expensive], ) _, provider_map, unique_models = create_model_mappings( diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py index 2d87a8cf..7f90b4f1 100644 --- a/tests/unit/test_certification_cache.py +++ b/tests/unit/test_certification_cache.py @@ -11,7 +11,12 @@ import pytest import respx from routstr.payment.cost_calculation import CostData, CostDataError -from routstr.upstream.certification import STATUS_FAIL, STATUS_OK, STATUS_WARN +from routstr.upstream.certification import ( + PROBE_MAX_TOKENS, + STATUS_FAIL, + STATUS_OK, + STATUS_WARN, +) from routstr.upstream.certification_cache import ( CacheProbeResult, _raw_cache_keys, @@ -210,6 +215,8 @@ class TestCacheBillingRow: model=_model(), probe=_probe([_payload(UNCACHED), cached]), cost_data=_cost(100), + provider_fee=1.0, + sats_to_usd=SATS_USD, ) assert row["status"] == STATUS_OK assert row["evidence"]["reported_usd"] == 5e-5 @@ -242,12 +249,21 @@ class TestCostMarginRow: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 1e-3}) row = self._row([payload]) assert row["status"] == STATUS_FAIL - assert "lose money" in row["detail"] + assert "below raw upstream cost" in row["detail"] + assert "does not measure client debits" in row["detail"] def test_fee_scales_upstream_cost(self) -> None: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) assert self._row([payload], fee=1.0)["status"] == STATUS_OK - assert self._row([payload], fee=2.0)["status"] == STATUS_FAIL + marked_up = self._row([payload], fee=2.0) + assert marked_up["status"] == STATUS_FAIL + assert "Raw upstream cost is covered" in marked_up["detail"] + assert "lose money" not in marked_up["detail"] + sample = marked_up["evidence"]["samples"][0] + assert sample["upstream_msats"] == 2 + assert sample["upstream_msats_with_fee"] == 4 + assert sample["token_estimate_minus_upstream_msats"] == 0 + assert sample["token_estimate_minus_fee_target_msats"] == -2 def test_deepseek_endpoint_price_exceeds_configured_model_price(self) -> None: sats_usd = 0.0008616302499999999 @@ -332,6 +348,8 @@ class TestCostMarginRow: ) assert row["status"] == STATUS_OK + assert row["title"] == "Token estimate covers fee-adjusted target" + assert "raw upstream cost exceeds" in row["detail"] def test_warn_when_pricing_unknown(self) -> None: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) @@ -373,7 +391,7 @@ class TestProbeCache: assert first == second system = first["messages"][0]["content"] assert system[0]["cache_control"] == {"type": "ephemeral"} - assert first["max_tokens"] == 1 + assert first["max_tokens"] == PROBE_MAX_TOKENS assert route.calls[0].request.headers["Authorization"] == "Bearer k" @pytest.mark.asyncio diff --git a/tests/unit/test_certification_deep_review.py b/tests/unit/test_certification_deep_review.py new file mode 100644 index 00000000..6ffad8de --- /dev/null +++ b/tests/unit/test_certification_deep_review.py @@ -0,0 +1,419 @@ +"""Regressions from the independent three-pass certification audit.""" + +import json +from typing import Any + +import httpx +import pytest + +from routstr.payment import price +from routstr.payment.cost_calculation import CostData +from routstr.upstream.certification import ( + ProbeResult, + _model_from_usd_pricing, + build_checklist, + cost_prompt_completion_row, + run_live_checks, + usage_capture_row, +) +from routstr.upstream.certification_cache import ( + CacheProbeResult, + cache_billing_row, + cache_reported_row, + probe_cache, +) +from routstr.upstream.certification_probe import ( + PROBE_TOKEN_BUDGETS, + next_probe_budget, +) + +SATS_USD = 0.0005 +URL = "https://mock.example/v1" +CACHED = { + "prompt_tokens": 50, + "completion_tokens": 0, + "prompt_tokens_details": {"cached_tokens": 45}, + "cost": 5e-5, +} +EXHAUSTED = { + "error": { + "message": "Could not finish the message because max_tokens " + "or model output limit was reached. Please try again with higher max_tokens." + } +} + + +def model(name: str = "test-model") -> Any: + return _model_from_usd_pricing(name, 1e-6, 2e-6, SATS_USD) + + +def cost(total: int) -> CostData: + return CostData(base_msats=0, input_msats=total, output_msats=0, total_msats=total) + + +def cached_probe(usage: dict[str, Any]) -> CacheProbeResult: + return CacheProbeResult( + chat_url=f"{URL}/chat/completions", + statuses=[200, 200], + payloads=[{"usage": usage}, {"usage": usage}], + ) + + +def rows_by_id(rows: list[dict[str, Any]]) -> dict[str, Any]: + return {row["id"]: row for row in rows} + + +@pytest.fixture(autouse=True) +def local_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(price, "SATS_USD_PRICE", SATS_USD) + monkeypatch.setattr(settings, "fixed_pricing", False) + + +@pytest.mark.parametrize("wrong", [0, 1, 999999]) +def test_cached_usd_rejects_wrong_calculation(wrong: int) -> None: + row = cache_billing_row( + model=model(), + probe=cached_probe(CACHED), + cost_data=cost(wrong), + provider_fee=1.25, + sats_to_usd=SATS_USD, + ) + assert row["status"] == "fail" + assert row["evidence"]["basis"] == "upstream_reported_usd" + assert row["evidence"]["expected_total_msats"] == 125 + + +def test_cached_usd_verifies_fee_without_claiming_discount() -> None: + row = cache_billing_row( + model=model(), + probe=cached_probe(CACHED), + cost_data=cost(125), + provider_fee=1.25, + sats_to_usd=SATS_USD, + ) + assert row["status"] == "ok" + assert row["evidence"]["expected_total_msats"] == 125 + assert row["evidence"]["upstream_discount_verified"] is False + assert "discount" not in row["title"].lower() + assert "not verified" in row["detail"] + + +def test_equal_full_price_does_not_certify_discount() -> None: + row = cache_billing_row( + model=model(), + probe=cached_probe(CACHED), + cost_data=cost(100), + sats_to_usd=SATS_USD, + ) + assert row["status"] == "ok" + assert ( + row["evidence"]["actual_total_msats"] + == row["evidence"]["full_price_total_msats"] + ) + assert row["evidence"]["upstream_discount_verified"] is False + assert "already carries" not in row["detail"] + assert "cache rate" not in build_checklist([row])[4]["label"] + + +@pytest.mark.parametrize( + "rate,status", [(None, "warn"), (0.0, "fail"), (float("nan"), "fail")] +) +def test_cached_usd_requires_valid_exchange_rate( + rate: float | None, status: str +) -> None: + row = cache_billing_row( + model=model(), probe=cached_probe(CACHED), cost_data=cost(100), sats_to_usd=rate + ) + assert row["status"] == status + + +@pytest.mark.parametrize("free", [False, True]) +def test_unknown_cost_does_not_claim_upstream_overpayment(free: bool) -> None: + usage = {key: value for key, value in CACHED.items() if key != "cost"} + priced = _model_from_usd_pricing("test-model", 0 if free else 1e-6, 0, SATS_USD) + row = cache_billing_row( + model=priced, probe=cached_probe(usage), cost_data=cost(0 if free else 100) + ) + assert row["status"] == "warn" + assert "pay more" not in row["detail"] + assert "unverified" in row["detail"] + + +def test_missing_usage_does_not_claim_free_settlement() -> None: + probe = ProbeResult( + base_url=URL, + models_url=f"{URL}/models", + chat_url=f"{URL}/chat/completions", + chat_status=200, + chat_payload={"choices": [{"message": {"content": "ok"}}]}, + ) + row = usage_capture_row(probe) + assert row["status"] == "warn" + assert "estimate" in row["detail"] + assert "(0+0)" not in row["detail"] + + +def test_short_usd_basis_is_not_called_configured_pricing() -> None: + probe = ProbeResult( + base_url=URL, + models_url=f"{URL}/models", + chat_url=f"{URL}/chat/completions", + chat_status=200, + chat_payload={"usage": CACHED}, + ) + row = cost_prompt_completion_row( + model=model(), + probe=probe, + cost_data=cost(100), + provider_fee=1, + sats_to_usd=SATS_USD, + ) + assert row["status"] == "ok" + assert "upstream-reported USD" in row["detail"] + assert "matching the configured" not in row["detail"] + + +def test_no_evaluated_models_warns_all_pricing_rows() -> None: + from routstr.core.admin import ( + _report_row_cache_rate, + _report_row_enabled_models_served, + _report_row_sats_pricing_present, + _report_row_served_matches_configured, + ) + + rows = [ + builder([]) + for builder in ( + _report_row_cache_rate, + _report_row_enabled_models_served, + _report_row_sats_pricing_present, + _report_row_served_matches_configured, + ) + ] + assert all( + row["status"] == "warn" and row["evidence"]["checked"] == 0 for row in rows + ) + assert all("Not evaluated" in row["detail"] for row in rows) + assert ( + next(c for c in build_checklist(rows) if c["goal"] == "pricing_v1_models")[ + "status" + ] + == "warn" + ) + + +@pytest.mark.asyncio +async def test_typesafe_native_listing_without_unsupported_chat() -> None: + from routstr.upstream.typesafe import TypeSafeUpstreamProvider + + calls = [] + + def handle(request: httpx.Request) -> httpx.Response: + calls.append(request) + assert request.method == "GET" + return httpx.Response(200, json={"models": [{"name": "jev-1.13"}]}) + + upstream = TypeSafeUpstreamProvider(api_key="test-key") + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + rows = rows_by_id( + await run_live_checks( + upstream.base_url, + "test-key", + model("jev-1.13.0"), + provider_fee=1, + sats_to_usd=SATS_USD, + upstream=upstream, + client=client, + ) + ) + assert len(calls) == 1 + assert rows["endpoint.models_payload"]["status"] == "ok" + assert rows["endpoint.models_payload"]["evidence"]["sample_ids"] == ["jev-1.13"] + for key in ( + "usage.capture", + "cost.prompt_completion", + "cache.reported", + "cache.billing", + "cost.margin", + ): + assert rows[key]["status"] == "warn" + assert "unverified" in rows[key]["detail"] + + +@pytest.mark.asyncio +async def test_corrected_field_and_budget_propagate_to_cache() -> None: + bodies = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json={"data": [{"id": "test-model"}]}) + body = json.loads(request.content) + bodies.append(body) + if "max_tokens" in body: + return httpx.Response( + 400, + json={ + "error": { + "message": "Unsupported parameter: 'max_tokens'. Use 'max_completion_tokens' instead." + } + }, + ) + if body["max_completion_tokens"] < 128: + return httpx.Response( + 400, + json={ + "error": { + "message": "max_completion_tokens must be at least 128 for this deployment." + } + }, + ) + return httpx.Response(200, json={"usage": CACHED}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + rows = rows_by_id( + await run_live_checks( + URL, "", model(), provider_fee=1, sats_to_usd=SATS_USD, client=client + ) + ) + assert rows["usage.capture"]["status"] == "ok" + assert len(bodies) == 5 + assert [body.get("max_completion_tokens") for body in bodies] == [ + None, + 32, + 128, + 128, + 128, + ] + evidence = rows["usage.capture"]["evidence"] + assert evidence["token_limit_field"] == "max_completion_tokens" + assert evidence["token_limit"] == 128 + assert [a["status_code"] for a in evidence["attempts"]] == [400, 400, 200] + assert rows["cache.reported"]["evidence"]["token_limit"] == 128 + + +@pytest.mark.asyncio +async def test_openai_provider_shaped_field_is_recorded( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.upstream.openai import OpenAIUpstreamProvider + + monkeypatch.setattr("routstr.upstream.openai._rejects_max_tokens", lambda _: True) + + def handle(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json={"data": [{"id": "gpt-5"}]}) + body = json.loads(request.content) + assert body["max_completion_tokens"] == 32 + assert "max_tokens" not in body + return httpx.Response(200, json={"usage": CACHED}) + + upstream = OpenAIUpstreamProvider(api_key="test-key") + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + rows = rows_by_id( + await run_live_checks( + upstream.base_url, + "test-key", + model("gpt-5"), + provider_fee=1, + sats_to_usd=SATS_USD, + upstream=upstream, + client=client, + ) + ) + assert ( + rows["usage.capture"]["evidence"]["token_limit_field"] + == "max_completion_tokens" + ) + + +@pytest.mark.asyncio +async def test_repeated_exhaustion_stops_at_ceiling_and_preserves_failure() -> None: + budgets = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.method == "GET": + return httpx.Response(200, json={"data": [{"id": "test-model"}]}) + budgets.append(json.loads(request.content)["max_tokens"]) + return httpx.Response(400, json=EXHAUSTED) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + rows = rows_by_id( + await run_live_checks( + URL, "", model(), provider_fee=1, sats_to_usd=SATS_USD, client=client + ) + ) + assert budgets == list(PROBE_TOKEN_BUDGETS) + assert rows["usage.capture"]["status"] == "fail" + assert len(rows["usage.capture"]["evidence"]["attempts"]) == len( + PROBE_TOKEN_BUDGETS + ) + assert rows["cache.reported"]["status"] == "warn" + + +@pytest.mark.parametrize( + "status,message,current,expected", + [ + (400, "max_tokens must be at least 16 for this deployment", 1, 32), + (400, "max_tokens must be at least 2 for this deployment", 1, 32), + (400, "max_tokens must be at least 9000", 32, None), + (400, "context window exceeded", 32, None), + (400, "unknown model", 32, None), + (429, "increase max_tokens", 32, None), + ], +) +def test_budget_corrections_are_narrow_and_bounded( + status: int, message: str, current: int, expected: int | None +) -> None: + assert ( + next_probe_budget( + status, + {"error": {"metadata": {"raw": json.dumps({"message": message})}}}, + current, + ) + == expected + ) + + +@pytest.mark.asyncio +async def test_cache_retains_fallback_and_http_error_without_credentials() -> None: + replies = iter( + [ + httpx.Response( + 400, json={"error": "cache_control not allowed; key=private-test-key"} + ), + httpx.Response(200, json={"usage": CACHED}), + httpx.Response(429, json={"error": "rate limited; key=private-test-key"}), + ] + ) + async with httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: next(replies)) + ) as client: + probe = await probe_cache(URL, "private-test-key", "test-model", client=client) + row = cache_reported_row(probe) + assert probe.statuses == [200, 429] + assert [a["status_code"] for a in row["evidence"]["attempts"]] == [400, 200, 429] + assert [a["request_format"] for a in probe.attempts] == [ + "cache_control", + "plain", + "plain", + ] + assert "cache_control not allowed" in probe.attempts[0]["body"] + assert "HTTP 429" in row["detail"] and "rate limited" in row["detail"] + assert "no response body" not in row["detail"] + assert "private-test-key" not in json.dumps(row) + + +@pytest.mark.asyncio +async def test_cache_decode_failure_has_bounded_body() -> None: + async with httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response(200, text="not JSON " + "x" * 3000) + ) + ) as client: + probe = await probe_cache(URL, "", "test-model", client=client) + row = cache_reported_row(probe) + assert row["status"] == "fail" + assert "JSONDecodeError" in row["detail"] + assert all(len(a["body"]) <= 1001 for a in probe.attempts) diff --git a/tests/unit/test_certification_token_limit.py b/tests/unit/test_certification_token_limit.py index 45c8b4da..adc3eb5c 100644 --- a/tests/unit/test_certification_token_limit.py +++ b/tests/unit/test_certification_token_limit.py @@ -13,6 +13,7 @@ import httpx import pytest from routstr.upstream.certification import ( + PROBE_MAX_TOKENS, STATUS_OK, certify_upstream_url, wants_max_completion_tokens, @@ -79,7 +80,9 @@ async def test_probe_retries_with_max_completion_tokens() -> None: # Rejected probe, retried probe, then both cache-probe calls reuse the # accepted field instead of being rejected again. assert ["max_tokens" in body for body in bodies] == [True, False, False, False] - assert all(body.get("max_completion_tokens") == 1 for body in bodies[1:]) + assert all( + body.get("max_completion_tokens") == PROBE_MAX_TOKENS for body in bodies[1:] + ) @pytest.mark.asyncio diff --git a/tests/unit/test_native_cache_pricing_policy.py b/tests/unit/test_native_cache_pricing_policy.py new file mode 100644 index 00000000..63801cc3 --- /dev/null +++ b/tests/unit/test_native_cache_pricing_policy.py @@ -0,0 +1,131 @@ +import json +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from routstr.core.db import ModelRow +from routstr.payment.models import Pricing, _row_to_model, list_models +from routstr.upstream.helpers import get_all_models_with_overrides +from routstr.upstream.model_paths import _price_in_sats + + +def _row(cache_rate: float = 0) -> ModelRow: + return ModelRow( + id="vendor/model", + name="Model", + created=0, + description="", + context_length=8192, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + } + ), + pricing=json.dumps( + {"prompt": 4e-6, "completion": 8e-6, "input_cache_read": cache_rate} + ), + enabled=True, + upstream_provider_id=1, + ) + + +def _session(row: ModelRow, provider_type: str): + provider = SimpleNamespace( + id=1, enabled=True, provider_fee=1.1, provider_type=provider_type + ) + return SimpleNamespace( + exec=AsyncMock( + side_effect=[ + SimpleNamespace(all=lambda: [row]), + SimpleNamespace(all=lambda: [provider]), + ] + ) + ) + + +@pytest.mark.parametrize("provider_type", ["ppqai", "venice"]) +@pytest.mark.parametrize("cache_rate", [0, 7e-7]) +def test_db_conversion_preserves_native_or_explicit_cache_prices( + provider_type, cache_rate +) -> None: + row = _row(cache_rate) + stored = row.pricing + with ( + patch( + "routstr.payment.models.backfill_cache_pricing", + return_value=Pricing(prompt=4e-6, completion=8e-6, input_cache_read=9e-7), + ) as backfill, + patch("routstr.payment.models.sats_usd_price", return_value=0.001), + ): + model = _row_to_model(row, True, 1.1, provider_type=provider_type) + backfill.assert_not_called() + assert model.pricing.input_cache_read == pytest.approx(cache_rate * 1.1) + assert row.pricing == stored + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider_type", ["ppqai", "venice"]) +async def test_admin_listing_passes_provider_policy_to_database_conversion( + provider_type, +) -> None: + with ( + patch("routstr.payment.models.backfill_cache_pricing") as backfill, + patch("routstr.payment.models.sats_usd_price", return_value=0.001), + ): + models = await list_models(_session(_row(), provider_type), upstream_id=1) + backfill.assert_not_called() + assert len(models) == 1 + assert models[0].pricing.input_cache_read == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider_type", ["ppqai", "venice"]) +async def test_runtime_overrides_do_not_reintroduce_generic_cache_prices( + provider_type, +) -> None: + row = _row() + session = _session(row, provider_type) + + @asynccontextmanager + async def create_session(): + yield session + + upstream = SimpleNamespace( + db_id=1, + provider_type=provider_type, + base_url="https://example.invalid", + get_cached_models=MagicMock( + return_value=[SimpleNamespace(id=row.id, enabled=True)] + ), + ) + with ( + patch("routstr.upstream.helpers.create_session", create_session), + patch("routstr.payment.models.backfill_cache_pricing") as backfill, + patch("routstr.payment.models.sats_usd_price", return_value=0.001), + ): + models = await get_all_models_with_overrides([upstream]) + backfill.assert_not_called() + assert len(models) == 1 + assert models[0].pricing.input_cache_read == 0 + + +@pytest.mark.parametrize("provider_type", ["ppqai", "venice"]) +def test_path_sats_conversion_respects_native_cache_policy(provider_type) -> None: + model = { + "id": "vendor/model", + "pricing": {"prompt": 4e-6, "completion": 8e-6}, + "context_length": 8192, + } + with ( + patch("routstr.payment.models.backfill_cache_pricing") as backfill, + patch("routstr.payment.price.sats_usd_price", return_value=0.001), + ): + _price_in_sats(model, 1.1, provider_type) + backfill.assert_not_called() + assert model["pricing"]["input_cache_read"] == 0 + assert model["sats_pricing"]["input_cache_read"] == 0 diff --git a/tests/unit/test_ppq_pricing_provenance.py b/tests/unit/test_ppq_pricing_provenance.py new file mode 100644 index 00000000..67ac3177 --- /dev/null +++ b/tests/unit/test_ppq_pricing_provenance.py @@ -0,0 +1,138 @@ +import json +from unittest.mock import AsyncMock, patch + +import httpx +import pytest + +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream.ppqai import PPQAIUpstreamProvider +from routstr.upstream.venice import VeniceUpstreamProvider + + +def _model() -> Model: + return Model( + id="vendor/model", + name="Model", + created=0, + description="Metadata", + context_length=8192, + architecture=Architecture( + modality="text->text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="Unknown", + ), + pricing=Pricing( + prompt=1e-6, + completion=2e-6, + input_cache_read=1e-8, + input_cache_write=3e-6, + request=0.1, + ), + ) + + +def _entry(model_id: str, pricing: dict) -> dict: + return { + "id": model_id, + "name": model_id, + "created_at": 0, + "context_length": 8192, + "pricing": pricing, + } + + +async def _fetch( + entries: list[dict], metadata: list[dict] | None = None +) -> list[Model]: + response = httpx.Response( + 200, + request=httpx.Request("GET", "https://api.ppq.ai/models"), + text=json.dumps({"data": entries}), + ) + with ( + patch( + "routstr.upstream.ppqai._safe_read_request", + AsyncMock(return_value=response), + ), + patch( + "routstr.upstream.ppqai.async_fetch_openrouter_models", + AsyncMock( + return_value=metadata if metadata is not None else [_model().dict()] + ), + ), + ): + return await PPQAIUpstreamProvider("test-only").fetch_models() + + +@pytest.mark.asyncio +async def test_ppq_does_not_inherit_other_provider_cache_or_request_rates() -> None: + (model,) = await _fetch( + [_entry("vendor/model", {"api": {"input_per_1M": 4, "output_per_1M": 8}})] + ) + assert model.pricing == Pricing(prompt=4e-6, completion=8e-6) + + +@pytest.mark.asyncio +async def test_ppq_alias_matches_do_not_share_mutated_prices() -> None: + models = await _fetch( + [ + _entry("vendor/model", {"api": {"input_per_1M": 4, "output_per_1M": 8}}), + _entry("model", {"api": {"input_per_1M": 6, "output_per_1M": 9}}), + ] + ) + assert [m.id for m in models] == ["vendor/model", "model"] + assert [m.pricing.prompt for m in models] == [4e-6, 6e-6] + assert models[0] is not models[1] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("metadata", [None, []]) +async def test_ppq_partial_api_prices_fall_back_per_field_preserving_zero( + metadata, +) -> None: + (model,) = await _fetch( + [ + _entry( + "vendor/model", + { + "api": {"input_per_1M": 0}, + "input_per_1M_tokens": 5, + "output_per_1M_tokens": 8, + }, + ) + ], + metadata, + ) + assert model.pricing.prompt == 0 + assert model.pricing.completion == 8e-6 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("rate", [None, -1, float("inf"), float("nan")]) +async def test_ppq_unpriced_or_invalid_native_rate_is_not_replaced_by_openrouter( + rate, +) -> None: + assert ( + await _fetch( + [ + _entry( + "vendor/model", {"api": {"input_per_1M": rate, "output_per_1M": 8}} + ) + ] + ) + == [] + ) + + +@pytest.mark.parametrize("provider", [PPQAIUpstreamProvider, VeniceUpstreamProvider]) +def test_native_catalog_providers_do_not_backfill_generic_cache_rates(provider) -> None: + model = _model() + model.pricing = Pricing(prompt=4e-6, completion=8e-6) + with patch("routstr.upstream.base.backfill_cache_pricing") as backfill: + adjusted = provider("test-only", provider_fee=1.1)._apply_provider_fee_to_model( + model + ) + backfill.assert_not_called() + assert adjusted.pricing.input_cache_read == 0 + assert adjusted.pricing.prompt == pytest.approx(4.4e-6)