mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: cover certification review leftovers and shape model tests like the proxy
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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={
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user