fix: cover certification review leftovers and shape model tests like the proxy

This commit is contained in:
9qeklajc
2026-10-02 20:05:39 +02:00
parent b99f644b81
commit a0aad288aa
12 changed files with 436 additions and 117 deletions
+6 -1
View File
@@ -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
+45 -11
View File
@@ -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:
+4 -13
View File
@@ -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
+89 -39
View File
@@ -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
+51 -11
View File
@@ -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,
),
),
]
+2 -2
View File
@@ -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
+13 -5
View File
@@ -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
@@ -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
@@ -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={
+41
View File
@@ -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(
+6 -33
View File
@@ -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
@@ -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