From a0aad288aac8a4254299c27e38af94ae562ae449 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 2 Oct 2026 20:05:39 +0200 Subject: [PATCH] fix: cover certification review leftovers and shape model tests like the proxy --- routstr/core/admin.py | 7 +- routstr/payment/models.py | 56 ++++++-- routstr/payment/usage.py | 17 +-- routstr/upstream/certification.py | 128 ++++++++++++------ routstr/upstream/certification_cache.py | 62 +++++++-- routstr/upstream/model_paths.py | 4 +- tests/integration/test_certify_endpoint.py | 18 ++- .../test_certify_provider_shapes.py | 46 +++++++ .../test_model_test_endpoint_security.py | 12 +- tests/unit/test_certification_cache.py | 41 ++++++ tests/unit/test_certification_hardening.py | 39 +----- tests/unit/test_certification_token_limit.py | 123 +++++++++++++++++ 12 files changed, 436 insertions(+), 117 deletions(-) create mode 100644 tests/unit/test_certification_token_limit.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 2027cc16..957684af 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1720,10 +1720,14 @@ async def certify_upstream_provider( status_code=503, detail="sats/USD price is not initialized yet; retry shortly", ) + # The proxy reserves and token-bills a pinned request with the model's + # own pricing, so the cost rows use it too; the path's advertised + # endpoint rates are only compared against it in the margin row. + advertised_model = None if selected_path is not None: from ..upstream.model_paths import apply_model_path_pricing - model_obj = apply_model_path_pricing( + advertised_model = apply_model_path_pricing( model_obj, selected_path, provider.provider_fee, @@ -1741,6 +1745,7 @@ async def certify_upstream_provider( check_cache=payload.check_cache, endpoint_tag=endpoint_tag, upstream=upstream_obj, + advertised_model=advertised_model, ) rows = pricing_rows + live_rows diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 8d7788a1..731c6de5 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -655,6 +655,44 @@ class ModelTestRequest(V2BaseModel): request_data: dict +def _model_test_target( + provider: UpstreamProviderRow, + model_row: ModelRow, + endpoint_path: str, + model_id: str, +) -> tuple[str, dict[str, str], dict[str, str], str]: + """URL, headers, query params and model id for a model test, shaped like + the proxy's. + + With the provider's live upstream instance, use the hooks + ``forward_request`` uses (Azure's deployment path, ``api-key`` and + ``api-version``, Gemini's ``/openai`` base, Ollama's ``/v1``, model-name + transforms). Without one, assume a plain OpenAI-compatible base URL. + """ + from ..proxy import get_upstreams + + upstream = next( + (u for u in get_upstreams() if getattr(u, "db_id", None) == provider.id), + None, + ) + if upstream is None: + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {provider.api_key}", + } + url = f"{provider.base_url.rstrip('/')}/{endpoint_path}" + return url, headers, {}, model_id + + model_obj = _build_model_from_row(model_row, False, provider.provider_fee) + path = upstream.normalize_request_path(f"v1/{endpoint_path}", model_obj) + return ( + upstream.build_request_url(path, model_obj), + upstream.prepare_headers({"content-type": "application/json"}), + dict(upstream.prepare_params(path, None)), + upstream.transform_model_name(model_id), + ) + + @models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) async def test_model( payload: ModelTestRequest, @@ -688,8 +726,10 @@ async def test_model( raise HTTPException(status_code=400, detail="Unsupported endpoint_type") actual_model_id = model_row.forwarded_model_id or model_row.id - request_data = dict(payload.request_data) - request_data["model"] = actual_model_id + url, headers, params, upstream_model_id = _model_test_target( + provider, model_row, endpoint_path, actual_model_id + ) + request_data = {**payload.request_data, "model": upstream_model_id} try: request_size = len(json.dumps(request_data).encode("utf-8")) @@ -698,9 +738,6 @@ async def test_model( if request_size > _MODEL_TEST_MAX_REQUEST_BYTES: raise HTTPException(status_code=413, detail="request_data too large") - base_url = provider.base_url.rstrip("/") - url = f"{base_url}/{endpoint_path}" - logger.info( "admin model test", extra={ @@ -712,14 +749,11 @@ async def test_model( }, ) - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {provider.api_key}", - } - try: async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post(url, json=request_data, headers=headers) + response = await client.post( + url, json=request_data, headers=headers, params=params + ) try: response_data = response.json() except Exception: diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index d69ab55c..08d675ad 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -38,8 +38,6 @@ names do not collide, so a single union parser is safe; a vendor whose fields would genuinely conflict needs a dedicated branch here. """ -import math - from pydantic.v1 import BaseModel @@ -53,25 +51,18 @@ class NormalizedUsage(BaseModel): def parse_token_count(value: object) -> int: - """Parse a token count from various formats (int, float, str, bool). - - ``json.loads`` accepts bare ``Infinity``/``NaN`` and overflows ``1e999`` to - ``inf``, so an upstream can put them on the wire. ``int()`` raises on both, - which would turn a billing path into a 500; reject them like - ``is_usable_rate`` does instead. - """ + """Parse a token count from various formats (int, float, str, bool).""" if isinstance(value, bool): return 0 if isinstance(value, int): return max(0, value) if isinstance(value, float): - return max(0, int(value)) if math.isfinite(value) else 0 + return max(0, int(value)) if isinstance(value, str): try: - parsed = float(value) - except (ValueError, OverflowError): + return max(0, int(float(value))) + except ValueError: return 0 - return max(0, int(parsed)) if math.isfinite(parsed) else 0 return 0 diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index 2989e287..cb4d83e4 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -179,6 +179,19 @@ 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) @dataclass @@ -260,8 +273,10 @@ async def probe_upstream( ) -> ProbeResult: """Call the upstream's ``/models`` and a one-token completion. - Each HTTP call, including its body read, has an elapsed-time deadline. - A transport failure is a ``fail`` row, not a failed admin request. + 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. """ shape = probe_shape(base_url, api_key, upstream, model) result = ProbeResult( @@ -300,44 +315,22 @@ async def probe_upstream( result.models_error = f"{type(exc).__name__}: {exc}" result.models_latency_ms = round((time.monotonic() - started) * 1000, 2) - started = time.monotonic() if not model_id: return result - request_body = { - "model": model_id, - "messages": [{"role": "user", "content": PROBE_PROMPT}], - "max_tokens": PROBE_MAX_TOKENS, - "stream": False, - } - if endpoint_tag: - request_body["provider"] = { - "order": [endpoint_tag], - "allow_fallbacks": False, - } - try: - async with asyncio.timeout(timeout): - response = await client.post( - result.chat_url, - json=shape_body(request_body, upstream, model), - headers=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) + 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() @@ -345,6 +338,59 @@ async def probe_upstream( return result +async def _probe_chat( + client: httpx.AsyncClient, + result: ProbeResult, + model_id: str, + shape: ProbeShape, + timeout: float, + upstream: "BaseUpstreamProvider | None", + model: "Model | None", + token_field: str, +) -> None: + """Send the one-token completion and record its outcome on ``result``.""" + request_body: dict[str, Any] = { + "model": model_id, + "messages": [{"role": "user", "content": PROBE_PROMPT}], + token_field: PROBE_MAX_TOKENS, + "stream": False, + } + if result.endpoint_tag: + request_body["provider"] = { + "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) + + # Row builders are pure: the network lives only in ``probe_upstream`` and # ``run_live_checks``, so every verdict is testable without a socket. @@ -817,13 +863,15 @@ async def run_live_checks( check_cache: bool = True, endpoint_tag: str | None = None, upstream: "BaseUpstreamProvider | None" = None, + advertised_model: "Model | None" = None, ) -> list[dict[str, Any]]: """Probe one upstream and build the live/derived rows. ``check_cache`` adds the prompt-cache and margin rows, which cost two or three more completions against a long prompt. ``upstream`` shapes the probes like the proxy's own requests; without it they assume a plain - OpenAI-compatible base URL. + OpenAI-compatible base URL. ``advertised_model`` carries a pinned path's + own endpoint rates for the margin row to compare against ``model``'s. """ probe = await probe_upstream( base_url, @@ -927,6 +975,8 @@ async def run_live_checks( pricing_known=pricing_known, endpoint_tag=endpoint_tag, upstream=upstream, + advertised_model=advertised_model, + token_limit_field=probe.token_limit_field, ) ) return rows diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index 9c267b69..a7459683 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -99,7 +99,11 @@ class CacheProbeResult: def _request_body( - model_id: str, prefix: str, fmt: str, endpoint_tag: str | None + model_id: str, + prefix: str, + fmt: str, + endpoint_tag: str | None, + token_field: str = "max_tokens", ) -> dict[str, Any]: system: Any if fmt == "cache_control": @@ -118,7 +122,7 @@ def _request_body( {"role": "system", "content": system}, {"role": "user", "content": CACHE_PROBE_QUESTION}, ], - "max_tokens": PROBE_MAX_TOKENS, + token_field: PROBE_MAX_TOKENS, "stream": False, } if endpoint_tag: @@ -184,6 +188,7 @@ async def probe_cache( timeout: float = PROBE_TIMEOUT_SECONDS, upstream: "BaseUpstreamProvider | None" = None, model: "Model | None" = None, + token_limit_field: str = "max_tokens", ) -> CacheProbeResult: """Send the same long prompt twice. @@ -202,7 +207,7 @@ async def probe_cache( async def post( fmt: str, ) -> tuple[int | None, dict[str, Any] | None, str | None, float]: - body = _request_body(model_id, prefix, fmt, endpoint_tag) + body = _request_body(model_id, prefix, fmt, endpoint_tag, token_limit_field) return await _post_completion( http, result.chat_url, @@ -441,6 +446,7 @@ def cost_margin_row( provider_fee: float, sats_to_usd: float, pricing_known: bool = True, + advertised_model: Model | None = None, ) -> dict[str, Any]: """Configured token pricing must cover what the upstream reports charging. @@ -450,11 +456,20 @@ def cost_margin_row( 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. + + ``model`` carries the pricing the proxy reserves and token-bills with. On + a pinned path, ``advertised_model`` carries the endpoint's own rates; a + covered margin whose advertised rates differ from the billed ones is a + warn, since ``/v1/models/paths`` then shows a price the node does not bill. """ + advertised_pricing = ( + advertised_model.sats_pricing if advertised_model is not None else None + ) evidence: dict[str, Any] = { "model_id": model.id, "provider_fee": provider_fee, "sats_usd_price": sats_to_usd, + "pricing_basis": "model pricing the proxy reserves and token-bills with", "samples": [], } if model.sats_pricing is None or not pricing_known: @@ -468,6 +483,7 @@ def cost_margin_row( samples: list[dict[str, Any]] = [] short: list[str] = [] + mismatched: list[str] = [] for payload in payloads: if not isinstance(payload, dict): continue @@ -480,6 +496,11 @@ def cost_margin_row( upstream_total = _expected_usd_msats( reported_usd, provider_fee, sats_to_usd ) + advertised_total = ( + _expected_token_msats(advertised_pricing, usage)[0] + if advertised_pricing is not None + else None + ) except (ValueError, OverflowError) as exc: evidence["error"] = f"{type(exc).__name__}: {exc}" return certification_row( @@ -489,14 +510,17 @@ def cost_margin_row( f"The margin could not be derived: {exc}.", evidence, ) - samples.append( - { - "usage": usage.dict(), - "reported_usd": reported_usd, - "upstream_msats_with_fee": upstream_total, - "configured_msats": configured_total, - } - ) + sample: dict[str, Any] = { + "usage": usage.dict(), + "reported_usd": reported_usd, + "upstream_msats_with_fee": upstream_total, + "configured_msats": configured_total, + } + if advertised_total is not None: + sample["advertised_msats"] = advertised_total + if abs(advertised_total - configured_total) > COST_TOLERANCE_MSATS: + mismatched.append(f"{advertised_total} vs {configured_total}") + samples.append(sample) if configured_total + COST_TOLERANCE_MSATS < upstream_total: short.append(f"{configured_total} < {upstream_total}") evidence["samples"] = samples @@ -521,6 +545,18 @@ def cost_margin_row( "requests lose money.", evidence, ) + if mismatched: + return certification_row( + ROW_MARGIN, + STATUS_WARN, + TITLE_MARGIN, + f"Configured pricing covers the upstream's reported cost on " + f"{len(samples)} sampled completion(s), but this path advertises " + f"different endpoint rates (advertised vs billed msats: " + f"{'; '.join(mismatched)}); the proxy reserves and token-bills " + "pinned requests with the model's own pricing.", + evidence, + ) return certification_row( ROW_MARGIN, STATUS_OK, @@ -561,6 +597,8 @@ async def run_cache_checks( pricing_known: bool = True, endpoint_tag: str | None = None, upstream: "BaseUpstreamProvider | None" = None, + advertised_model: Model | None = None, + token_limit_field: str = "max_tokens", ) -> list[dict[str, Any]]: """Run the cache probe and build the three cache/margin rows.""" probe = await probe_cache( @@ -572,6 +610,7 @@ async def run_cache_checks( timeout=timeout, upstream=upstream, model=model, + token_limit_field=token_limit_field, ) cost_data = await _price_payload(probe.second_payload, model, provider_fee) return [ @@ -595,6 +634,7 @@ async def run_cache_checks( provider_fee=provider_fee, sats_to_usd=sats_to_usd, pricing_known=pricing_known, + advertised_model=advertised_model, ), ), ] diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index db110e53..bcda8d39 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -898,8 +898,8 @@ def apply_model_path_pricing( Direct paths already use the provider model cache and therefore carry the same pricing as ``model``. OpenRouter endpoint rows instead contain raw, - endpoint-specific USD rates; certification must use those rates when its - requests are pinned to that endpoint. + endpoint-specific USD rates; certification compares them against the + model's own pricing, which the proxy reserves and token-bills with. """ if row.endpoint_tag is None: return model diff --git a/tests/integration/test_certify_endpoint.py b/tests/integration/test_certify_endpoint.py index a57ce706..9c8b2125 100644 --- a/tests/integration/test_certify_endpoint.py +++ b/tests/integration/test_certify_endpoint.py @@ -328,7 +328,7 @@ async def test_certify_model_path_pins_every_completion( @pytest.mark.integration @pytest.mark.asyncio @respx.mock -async def test_certify_uses_selected_path_pricing_for_margin( +async def test_certify_margin_bills_model_pricing_and_reports_path_pricing( integration_client: AsyncClient, integration_session: AsyncSession ) -> None: base_url = "https://openrouter.ai/api/v1" @@ -409,12 +409,20 @@ async def test_certify_uses_selected_path_pricing_for_margin( assert resp.status_code == 200, resp.text margin = _find_row(resp.json()["rows"], "cost.margin") + # The proxy reserves and token-bills a pinned request with the model's own + # pricing (``configured_msats``); the path's endpoint rates are reported + # alongside (``advertised_msats``) and differ, so the covered margin warns. assert [ - (sample["upstream_msats_with_fee"], sample["configured_msats"]) + ( + sample["upstream_msats_with_fee"], + sample["configured_msats"], + sample["advertised_msats"], + ) for sample in margin["evidence"]["samples"] - ] == [(3, 3), (269, 289), (26, 15)] - assert "289 < 269" not in margin["detail"] - assert "15 < 26" in margin["detail"] + ] == [(3, 3, 3), (269, 356, 289), (26, 43, 15)] + assert margin["status"] == "warn" + assert "advertises different endpoint rates" in margin["detail"] + assert "289 vs 356" in margin["detail"] @pytest.mark.integration diff --git a/tests/integration/test_certify_provider_shapes.py b/tests/integration/test_certify_provider_shapes.py index ee7ab153..c614bc26 100644 --- a/tests/integration/test_certify_provider_shapes.py +++ b/tests/integration/test_certify_provider_shapes.py @@ -88,6 +88,16 @@ SHAPES = [ ] +@pytest.fixture(autouse=True) +def _isolate_proxy_state(monkeypatch: pytest.MonkeyPatch) -> None: + """``reinitialize_upstreams`` rebinds module globals; restore them after + each test so the provider types seeded here never leak into others.""" + from routstr import proxy + + for name in ("_upstreams", "_provider_map", "_unique_models"): + monkeypatch.setattr(proxy, name, getattr(proxy, name).copy()) + + async def _seed(session: AsyncSession, shape: Shape) -> int: provider = UpstreamProviderRow( provider_type=shape.provider_type, @@ -183,3 +193,39 @@ def test_shape_body_keeps_a_single_cache_control_marker() -> None: assert json.dumps(shaped).count('"cache_control"') == 1 assert shaped["model"] == "claude-sonnet-4-5-20250929" + + +@pytest.mark.integration +@pytest.mark.asyncio +@pytest.mark.parametrize("shape", SHAPES, ids=[s.provider_type for s in SHAPES]) +async def test_model_test_matches_proxy_request( + shape: Shape, + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """``POST /api/models/test`` reaches the upstream the way the proxy does.""" + with respx.mock(assert_all_called=False) as mock: + await _seed(integration_session, shape) + chat_route = mock.post(shape.chat_url).mock( + return_value=Response(200, json=_mock_chat_response(model=shape.model_id)) + ) + + resp = await integration_client.post( + "/api/models/test", + headers=_admin_headers(), + json={ + "model_id": shape.model_id, + "endpoint_type": "chat-completions", + "request_data": {"messages": [{"role": "user", "content": "hi"}]}, + }, + ) + + assert resp.status_code == 200, resp.text + assert resp.json()["success"] is True, resp.json() + assert chat_route.call_count == 1 + request = chat_route.calls[0].request + header, value = shape.auth_header + assert request.headers.get(header) == value + for key, expected in shape.params.items(): + assert request.url.params.get(key) == expected + assert json.loads(request.content)["model"] == shape.upstream_model diff --git a/tests/integration/test_model_test_endpoint_security.py b/tests/integration/test_model_test_endpoint_security.py index dca7e9fa..b26a84e4 100644 --- a/tests/integration/test_model_test_endpoint_security.py +++ b/tests/integration/test_model_test_endpoint_security.py @@ -196,7 +196,11 @@ async def test_model_test_endpoint_admin_uses_allowed_upstream_path( return None async def post( - self, url: str, json: dict[str, Any], headers: dict[str, str] + self, + url: str, + json: dict[str, Any], + headers: dict[str, str], + params: dict[str, str] | None = None, ) -> MockResponse: assert url == "https://api.example.com/v1/chat/completions" assert json["model"] == "upstream-model-a" @@ -204,7 +208,11 @@ async def test_model_test_endpoint_admin_uses_allowed_upstream_path( return MockResponse() try: - with patch("httpx.AsyncClient", return_value=MockAsyncClient()): + # No live upstream instance: the plain OpenAI-compatible fallback. + with ( + patch("httpx.AsyncClient", return_value=MockAsyncClient()), + patch("routstr.proxy.get_upstreams", return_value=[]), + ): response = await integration_client.post( "/api/models/test", json={ diff --git a/tests/unit/test_certification_cache.py b/tests/unit/test_certification_cache.py index 2d87a8cf..a1681aeb 100644 --- a/tests/unit/test_certification_cache.py +++ b/tests/unit/test_certification_cache.py @@ -333,6 +333,47 @@ class TestCostMarginRow: assert row["status"] == STATUS_OK + def test_pinned_path_fails_when_billed_pricing_misses_cost(self) -> None: + """The path's endpoint rates cover the cost but the model pricing the + proxy actually bills with does not: the margin must fail.""" + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 4e-6}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + advertised_model=_model(prompt=1e-6, completion=2e-6), + ) + assert row["status"] == STATUS_FAIL + sample = row["evidence"]["samples"][0] + assert sample["advertised_msats"] >= sample["upstream_msats_with_fee"] + assert sample["configured_msats"] < sample["upstream_msats_with_fee"] + + def test_pinned_path_warns_when_advertised_rates_differ(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + advertised_model=_model(prompt=1e-6, completion=2e-6), + ) + assert row["status"] == STATUS_WARN + assert "advertises different endpoint rates" in row["detail"] + + def test_pinned_path_ok_when_advertised_rates_match(self) -> None: + payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) + row = cost_margin_row( + model=_model(), + payloads=[payload], + provider_fee=1.0, + sats_to_usd=SATS_USD, + advertised_model=_model(), + ) + assert row["status"] == STATUS_OK + sample = row["evidence"]["samples"][0] + assert sample["advertised_msats"] == sample["configured_msats"] + def test_warn_when_pricing_unknown(self) -> None: payload = _payload({"prompt_tokens": 5, "completion_tokens": 1, "cost": 9e-7}) row = cost_margin_row( diff --git a/tests/unit/test_certification_hardening.py b/tests/unit/test_certification_hardening.py index 7eb05552..b284748c 100644 --- a/tests/unit/test_certification_hardening.py +++ b/tests/unit/test_certification_hardening.py @@ -17,7 +17,6 @@ from typing import Any import pytest -from routstr.payment.usage import parse_token_count from routstr.upstream.certification import ( STATUS_FAIL, STATUS_OK, @@ -50,39 +49,14 @@ def _probe(**kwargs: Any) -> ProbeResult: ) -# Regression: a non-finite token count crashed the billing path. +# Regression: a non-finite token count must not crash the usage row. # ``json.loads`` accepts the bare ``Infinity``/``NaN`` literals, so an -# upstream can put them on the wire; ``int(inf)`` raised OverflowError and -# ``int(nan)`` raised ValueError inside ``parse_token_count``. +# upstream can put them on the wire. Whether the parser rejects them (``warn``, +# nothing to bill on) or raises (``fail``, unreadable usage), the row reports +# it instead of raising. class TestNonFiniteTokenCounts: - @pytest.mark.parametrize( - "value", - [ - float("inf"), - float("-inf"), - float("nan"), - 1e999, - "Infinity", - "NaN", - "-Infinity", - "1e999", - ], - ) - def test_parse_token_count_rejects_non_finite(self, value: Any) -> None: - assert parse_token_count(value) == 0 - - def test_parse_token_count_still_parses_ordinary_values(self) -> None: - assert parse_token_count(42) == 42 - assert parse_token_count("42") == 42 - assert parse_token_count(42.9) == 42 - assert parse_token_count("42.9") == 42 - assert parse_token_count(True) == 0 - assert parse_token_count(-5) == 0 - assert parse_token_count("not a number") == 0 - assert parse_token_count(None) == 0 - def test_usage_row_survives_infinite_tokens(self) -> None: row = usage_capture_row( _probe( @@ -95,8 +69,7 @@ class TestNonFiniteTokenCounts: }, ) ) - # Both counts collapse to 0, which is the "nothing to bill on" case. - assert row["status"] == STATUS_WARN + assert row["status"] in (STATUS_WARN, STATUS_FAIL) def test_usage_row_survives_infinite_tokens_in_a_string(self) -> None: row = usage_capture_row( @@ -105,7 +78,7 @@ class TestNonFiniteTokenCounts: chat_payload={"usage": {"prompt_tokens": "Infinity"}}, ) ) - assert row["status"] == STATUS_WARN + assert row["status"] in (STATUS_WARN, STATUS_FAIL) # Regression: ``certification_row`` stored non-dict evidence verbatim, so the diff --git a/tests/unit/test_certification_token_limit.py b/tests/unit/test_certification_token_limit.py new file mode 100644 index 00000000..45c8b4da --- /dev/null +++ b/tests/unit/test_certification_token_limit.py @@ -0,0 +1,123 @@ +"""Probes fall back to ``max_completion_tokens`` when ``max_tokens`` is refused. + +OpenAI's o-series and gpt-5 reject ``max_tokens`` on chat completions, so a +probe that only ever sends it fails ``usage.capture`` on a healthy upstream. +""" + +from __future__ import annotations + +import json +from typing import Any + +import httpx +import pytest + +from routstr.upstream.certification import ( + STATUS_OK, + certify_upstream_url, + wants_max_completion_tokens, +) + +OPENAI_REJECTION = { + "error": { + "message": ( + "Unsupported parameter: 'max_tokens' is not supported with this " + "model. Use 'max_completion_tokens' instead." + ), + "type": "invalid_request_error", + "param": "max_tokens", + "code": "unsupported_parameter", + } +} + + +@pytest.fixture(autouse=True) +def _restore_price_globals(monkeypatch: pytest.MonkeyPatch) -> None: + """``certify_upstream_url`` publishes ``sats_usd_price`` to the price + module's globals; restore them so no later test sees this quote.""" + from routstr.payment import price + + monkeypatch.setattr(price, "SATS_USD_PRICE", price.SATS_USD_PRICE) + monkeypatch.setattr(price, "BTC_USD_PRICE", price.BTC_USD_PRICE) + + +def _row(result: dict[str, Any], row_id: str) -> dict[str, Any]: + return next(row for row in result["rows"] if row["id"] == row_id) + + +@pytest.mark.asyncio +async def test_probe_retries_with_max_completion_tokens() -> None: + bodies: list[dict[str, Any]] = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "gpt-5"}]}) + body = json.loads(request.content) + bodies.append(body) + if "max_tokens" in body: + return httpx.Response(400, json=OPENAI_REJECTION) + return httpx.Response( + 200, + json={ + "model": "gpt-5", + "usage": {"prompt_tokens": 8, "completion_tokens": 1}, + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + result = await certify_upstream_url( + "https://mock.example/v1", + model_id="gpt-5", + prompt_price=1e-6, + completion_price=2e-6, + sats_usd_price=0.001, + client=client, + ) + + assert _row(result, "usage.capture")["status"] == STATUS_OK + assert _row(result, "cost.prompt_completion")["status"] == STATUS_OK + # 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:]) + + +@pytest.mark.asyncio +async def test_probe_does_not_retry_unrelated_400() -> None: + calls: list[dict[str, Any]] = [] + + def handle(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/models"): + return httpx.Response(200, json={"data": [{"id": "m"}]}) + calls.append(json.loads(request.content)) + return httpx.Response(400, json={"error": {"message": "model not found"}}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client: + result = await certify_upstream_url( + "https://mock.example/v1", + model_id="m", + prompt_price=1e-6, + completion_price=2e-6, + sats_usd_price=0.001, + client=client, + ) + + assert len(calls) == 1 + assert _row(result, "usage.capture")["status"] != STATUS_OK + + +@pytest.mark.parametrize( + ("status", "payload", "expected"), + [ + (400, OPENAI_REJECTION, True), + (400, {"error": {"message": "bad model"}}, False), + (422, OPENAI_REJECTION, False), + (200, OPENAI_REJECTION, False), + (400, None, False), + (None, None, False), + ], +) +def test_wants_max_completion_tokens( + status: int | None, payload: Any, expected: bool +) -> None: + assert wants_max_completion_tokens(status, payload) is expected