mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: shape certification probes through provider request hooks
This commit is contained in:
+26
-28
@@ -1523,7 +1523,6 @@ async def get_upstream_provider_report(provider_id: str) -> dict[str, object]:
|
||||
class CertifyRequest(BaseModel):
|
||||
model_id: str | None = None
|
||||
model_path: str | None = None
|
||||
timeout_seconds: float | None = None
|
||||
check_cache: bool = True
|
||||
|
||||
|
||||
@@ -1538,8 +1537,9 @@ async def certify_upstream_provider(
|
||||
|
||||
Unlike the read-only ``GET …/report``, this probes the upstream over the
|
||||
network and runs the node's cost engine on the real response. It never
|
||||
enters the billing path, so it costs at most one completion's worth of
|
||||
upstream credit and nothing from the node's wallet.
|
||||
enters the billing path, so it costs nothing from the node's wallet. Its
|
||||
upstream spend is a one-token completion, plus two or three one-token
|
||||
completions on a ~4.4k-token prompt when ``check_cache`` is set.
|
||||
|
||||
Returns the read-only report's four ``pricing.*`` rows (re-derived here so
|
||||
the certification is self-contained), the live rows from
|
||||
@@ -1547,7 +1547,6 @@ async def certify_upstream_provider(
|
||||
operator-facing goals.
|
||||
"""
|
||||
from ..upstream.certification import (
|
||||
MAX_PROBE_TIMEOUT_SECONDS,
|
||||
PROBE_TIMEOUT_SECONDS,
|
||||
build_checklist,
|
||||
run_live_checks,
|
||||
@@ -1565,6 +1564,7 @@ async def certify_upstream_provider(
|
||||
enabled_rows = list(result.all())
|
||||
|
||||
endpoint_tag: str | None = None
|
||||
path_model_id: str | None = None
|
||||
selected_path: ModelPathRow | None = None
|
||||
if payload.model_path is not None:
|
||||
from ..upstream.model_paths import decode_model_path
|
||||
@@ -1585,6 +1585,7 @@ async def certify_upstream_provider(
|
||||
detail="Model path is not available for this provider",
|
||||
)
|
||||
endpoint_tag = selector.endpoint_tag
|
||||
path_model_id = selector.model_id
|
||||
|
||||
evaluations = [
|
||||
_evaluate_model_row(row, provider, provider_pk) for row in enabled_rows
|
||||
@@ -1609,6 +1610,12 @@ async def certify_upstream_provider(
|
||||
|
||||
from ..proxy import get_candidates, get_upstreams
|
||||
|
||||
# The live upstream instance shapes the probes exactly like the proxy's
|
||||
# own requests (paths, auth headers, query params, model-name transforms).
|
||||
upstream_obj = next(
|
||||
(u for u in get_upstreams() if getattr(u, "db_id", None) == provider_pk),
|
||||
None,
|
||||
)
|
||||
model_obj = None
|
||||
if model_id:
|
||||
try:
|
||||
@@ -1633,20 +1640,15 @@ async def certify_upstream_provider(
|
||||
# The active upstream cache carries the same fee-adjusted USD and sats
|
||||
# pricing used by the proxy, so it is the authoritative fallback for a
|
||||
# pre-configuration certification probe.
|
||||
if model_obj is None:
|
||||
for upstream in get_upstreams():
|
||||
if getattr(upstream, "db_id", None) != provider_pk:
|
||||
continue
|
||||
model_obj = next(
|
||||
(
|
||||
model
|
||||
for model in upstream.get_cached_models()
|
||||
if model.id == model_id or model.forwarded_model_id == model_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if model_obj is not None:
|
||||
break
|
||||
if model_obj is None and upstream_obj is not None:
|
||||
model_obj = next(
|
||||
(
|
||||
model
|
||||
for model in upstream_obj.get_cached_models()
|
||||
if model.id == model_id or model.forwarded_model_id == model_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if selected_path is not None:
|
||||
from ..upstream.model_paths import exposed_model_id
|
||||
|
||||
@@ -1654,7 +1656,8 @@ async def certify_upstream_provider(
|
||||
if (
|
||||
payload.model_id is None
|
||||
or selected_id is None
|
||||
or selected_id.lower() != selector.model_id.lower()
|
||||
or path_model_id is None
|
||||
or selected_id.lower() != path_model_id.lower()
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
@@ -1726,23 +1729,18 @@ async def certify_upstream_provider(
|
||||
provider.provider_fee,
|
||||
sats_to_usd,
|
||||
)
|
||||
# Clamp the admin-supplied timeout per upstream call. The run makes up
|
||||
# to five calls, so the request can stay open for up to five times it.
|
||||
requested = (
|
||||
payload.timeout_seconds
|
||||
if payload.timeout_seconds is not None
|
||||
else PROBE_TIMEOUT_SECONDS
|
||||
)
|
||||
timeout = min(max(requested, 1.0), MAX_PROBE_TIMEOUT_SECONDS)
|
||||
# The timeout applies per upstream call. The run makes up to five
|
||||
# calls, so the request can stay open for up to five times it.
|
||||
live_rows = await run_live_checks(
|
||||
provider.base_url,
|
||||
provider.api_key,
|
||||
model_obj,
|
||||
provider_fee=provider.provider_fee,
|
||||
sats_to_usd=sats_to_usd,
|
||||
timeout=timeout,
|
||||
timeout=PROBE_TIMEOUT_SECONDS,
|
||||
check_cache=payload.check_cache,
|
||||
endpoint_tag=endpoint_tag,
|
||||
upstream=upstream_obj,
|
||||
)
|
||||
|
||||
rows = pricing_rows + live_rows
|
||||
|
||||
@@ -4,7 +4,9 @@ Extends the read-only pricing rows, which never touch the network, with the
|
||||
ones that must: a ``/models`` heartbeat and a one-token completion.
|
||||
|
||||
Probes call the upstream directly with ``httpx``, never through the node's
|
||||
billing path — no reservation, no Cashu, at most one token of upstream spend.
|
||||
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).
|
||||
They sit behind ``POST …/certify`` rather than the read-only ``GET …/report``
|
||||
because they can block for the length of the timeout.
|
||||
"""
|
||||
@@ -25,13 +27,14 @@ from urllib.parse import urlparse
|
||||
import httpx
|
||||
|
||||
from ..core.logging import get_logger
|
||||
from ..payment.cost_calculation import calculate_cost
|
||||
from ..payment.cost_calculation import _resolve_usd_cost, calculate_cost
|
||||
from ..payment.rates import coerce_rate
|
||||
from ..payment.usage import normalize_usage
|
||||
from .model_paths import is_openrouter_base_url
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..payment.models import Model
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -44,9 +47,6 @@ 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
|
||||
|
||||
# Ceiling for the caller-supplied timeout override.
|
||||
MAX_PROBE_TIMEOUT_SECONDS = 60.0
|
||||
|
||||
# The cheapest request that still exercises the usage/cost path.
|
||||
PROBE_MAX_TOKENS = 1
|
||||
PROBE_PROMPT = "ping"
|
||||
@@ -181,6 +181,72 @@ class ProbeResult:
|
||||
chat_latency_ms: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProbeShape:
|
||||
"""Where the probe calls go and how they are authenticated."""
|
||||
|
||||
models_url: str
|
||||
chat_url: str
|
||||
headers: dict[str, str]
|
||||
models_params: dict[str, str]
|
||||
chat_params: dict[str, str]
|
||||
|
||||
|
||||
def probe_shape(
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
model: "Model | None" = None,
|
||||
) -> ProbeShape:
|
||||
"""The URLs, headers and query params a probe sends.
|
||||
|
||||
With the node's upstream instance, use the hooks
|
||||
``BaseUpstreamProvider.forward_request`` uses, so the probe reaches what
|
||||
the proxy reaches (Azure's deployment path and ``api-key``, Gemini's
|
||||
``/openai`` base, Ollama's ``/v1``). Without one (the CLI), assume a plain
|
||||
OpenAI-compatible base URL.
|
||||
"""
|
||||
if upstream is None:
|
||||
base = base_url.rstrip("/")
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return ProbeShape(f"{base}/models", f"{base}/chat/completions", headers, {}, {})
|
||||
chat_path = upstream.normalize_request_path("v1/chat/completions", model)
|
||||
# Azure lists models under ``/openai/models``, not at its endpoint root.
|
||||
if upstream.provider_type == "azure":
|
||||
models_path = "openai/models"
|
||||
else:
|
||||
models_path = upstream.normalize_request_path("v1/models")
|
||||
return ProbeShape(
|
||||
models_url=upstream.build_request_url(models_path),
|
||||
chat_url=upstream.build_request_url(chat_path, model),
|
||||
headers=upstream.prepare_headers({"content-type": "application/json"}),
|
||||
models_params=dict(upstream.prepare_params(models_path, None)),
|
||||
chat_params=dict(upstream.prepare_params(chat_path, None)),
|
||||
)
|
||||
|
||||
|
||||
def shape_body(
|
||||
body: dict[str, Any],
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
model: "Model | None" = None,
|
||||
) -> Any:
|
||||
"""The JSON body the proxy would forward, model-name transforms included.
|
||||
|
||||
``prepare_request_body`` rewrites ``model`` from ``model.id``; the probe
|
||||
keeps the id it chose (``forwarded_model_id`` first) and only applies the
|
||||
provider's own name transform to it.
|
||||
"""
|
||||
if upstream is None or model is None:
|
||||
return body
|
||||
shaped = upstream.prepare_request_body(json.dumps(body).encode(), model)
|
||||
data = json.loads(shaped) if shaped else dict(body)
|
||||
if isinstance(data, dict) and isinstance(body.get("model"), str):
|
||||
data["model"] = upstream.transform_model_name(body["model"])
|
||||
return data
|
||||
|
||||
|
||||
async def probe_upstream(
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
@@ -189,22 +255,22 @@ async def probe_upstream(
|
||||
endpoint_tag: str | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
timeout: float = PROBE_TIMEOUT_SECONDS,
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
model: "Model | None" = None,
|
||||
) -> 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.
|
||||
"""
|
||||
base = base_url.rstrip("/")
|
||||
shape = probe_shape(base_url, api_key, upstream, model)
|
||||
result = ProbeResult(
|
||||
base_url=base_url,
|
||||
models_url=f"{base}/models",
|
||||
chat_url=f"{base}/chat/completions",
|
||||
models_url=shape.models_url,
|
||||
chat_url=shape.chat_url,
|
||||
endpoint_tag=endpoint_tag,
|
||||
)
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
headers = shape.headers
|
||||
|
||||
owns_client = client is None
|
||||
if client is None:
|
||||
@@ -214,7 +280,9 @@ async def probe_upstream(
|
||||
started = time.monotonic()
|
||||
try:
|
||||
async with asyncio.timeout(timeout):
|
||||
response = await client.get(result.models_url, headers=headers)
|
||||
response = await client.get(
|
||||
result.models_url, headers=headers, params=shape.models_params
|
||||
)
|
||||
result.models_status = response.status_code
|
||||
result.models_latency_ms = round((time.monotonic() - started) * 1000, 2)
|
||||
try:
|
||||
@@ -249,7 +317,10 @@ async def probe_upstream(
|
||||
try:
|
||||
async with asyncio.timeout(timeout):
|
||||
response = await client.post(
|
||||
result.chat_url, json=request_body, headers=headers
|
||||
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)
|
||||
@@ -497,27 +568,14 @@ def _truncate(value: Any, limit: int = 400) -> Any:
|
||||
def _reported_usd_cost(payload: dict[str, Any]) -> float:
|
||||
"""The upstream-reported USD cost, or 0.0 when it reported none.
|
||||
|
||||
Mirrors ``_resolve_usd_cost``'s priority and shares ``coerce_rate``, so
|
||||
this helper and the engine agree on *whether* a cost was reported; only
|
||||
the arithmetic below is re-derived independently.
|
||||
Uses the engine's own ``_resolve_usd_cost`` so both agree on *which*
|
||||
figure is the cost (PPQ.AI BYOK bills ``upstream_inference_cost`` plus the
|
||||
fee); only the arithmetic below is re-derived independently.
|
||||
"""
|
||||
usage = payload.get("usage")
|
||||
if not isinstance(usage, dict):
|
||||
return 0.0
|
||||
cost_details = usage.get("cost_details")
|
||||
if isinstance(cost_details, dict):
|
||||
total = coerce_rate(cost_details.get("total_cost"))
|
||||
if total is not None and total > 0:
|
||||
return total
|
||||
inference = coerce_rate(cost_details.get("upstream_inference_cost"))
|
||||
if inference is not None and inference > 0 and usage.get("is_byok"):
|
||||
return inference + (coerce_rate(usage.get("cost")) or 0.0)
|
||||
for source in (usage, payload):
|
||||
for field in ("total_cost", "cost"):
|
||||
value = coerce_rate(source.get(field))
|
||||
if value is not None and value > 0:
|
||||
return value
|
||||
return 0.0
|
||||
return _resolve_usd_cost(usage, payload)
|
||||
|
||||
|
||||
def _fixed_token_pricing_active() -> bool:
|
||||
@@ -758,11 +816,14 @@ async def run_live_checks(
|
||||
pricing_known: bool = True,
|
||||
check_cache: bool = True,
|
||||
endpoint_tag: str | None = None,
|
||||
upstream: "BaseUpstreamProvider | 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
|
||||
more completions against a long prompt.
|
||||
``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.
|
||||
"""
|
||||
probe = await probe_upstream(
|
||||
base_url,
|
||||
@@ -771,6 +832,8 @@ async def run_live_checks(
|
||||
endpoint_tag=endpoint_tag,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
upstream=upstream,
|
||||
model=model,
|
||||
)
|
||||
rows = [
|
||||
safe_row(
|
||||
@@ -863,6 +926,7 @@ async def run_live_checks(
|
||||
timeout=timeout,
|
||||
pricing_known=pricing_known,
|
||||
endpoint_tag=endpoint_tag,
|
||||
upstream=upstream,
|
||||
)
|
||||
)
|
||||
return rows
|
||||
|
||||
@@ -37,11 +37,14 @@ from .certification import (
|
||||
_reported_usd_cost,
|
||||
_token_rates,
|
||||
certification_row,
|
||||
probe_shape,
|
||||
safe_row,
|
||||
shape_body,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..payment.models import Model
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -129,14 +132,15 @@ def _request_body(
|
||||
async def _post_completion(
|
||||
client: httpx.AsyncClient,
|
||||
url: str,
|
||||
body: dict[str, Any],
|
||||
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)
|
||||
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
|
||||
@@ -178,6 +182,8 @@ async def probe_cache(
|
||||
endpoint_tag: str | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
timeout: float = PROBE_TIMEOUT_SECONDS,
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
model: "Model | None" = None,
|
||||
) -> CacheProbeResult:
|
||||
"""Send the same long prompt twice.
|
||||
|
||||
@@ -186,49 +192,39 @@ async def probe_cache(
|
||||
retry on HTTP 400/422, and the second call mirrors whichever format
|
||||
succeeded. Each call's elapsed deadline includes the response body.
|
||||
"""
|
||||
base = base_url.rstrip("/")
|
||||
result = CacheProbeResult(
|
||||
chat_url=f"{base}/chat/completions", endpoint_tag=endpoint_tag
|
||||
)
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
shape = probe_shape(base_url, api_key, upstream, model)
|
||||
result = CacheProbeResult(chat_url=shape.chat_url, endpoint_tag=endpoint_tag)
|
||||
prefix = cache_probe_prefix()
|
||||
|
||||
owns_client = client is None
|
||||
if client is None:
|
||||
client = httpx.AsyncClient(timeout=timeout)
|
||||
try:
|
||||
first = await _post_completion(
|
||||
client,
|
||||
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)
|
||||
return await _post_completion(
|
||||
http,
|
||||
result.chat_url,
|
||||
_request_body(model_id, prefix, "cache_control", endpoint_tag),
|
||||
headers,
|
||||
shape_body(body, upstream, model),
|
||||
shape.headers,
|
||||
shape.chat_params,
|
||||
timeout,
|
||||
)
|
||||
|
||||
try:
|
||||
first = await post("cache_control")
|
||||
if first[0] in (400, 422):
|
||||
result.request_format = "plain"
|
||||
first = await _post_completion(
|
||||
client,
|
||||
result.chat_url,
|
||||
_request_body(model_id, prefix, "plain", endpoint_tag),
|
||||
headers,
|
||||
timeout,
|
||||
)
|
||||
first = await post("plain")
|
||||
_record(result, first)
|
||||
if not _is_2xx(first[0]):
|
||||
return result
|
||||
second = await _post_completion(
|
||||
client,
|
||||
result.chat_url,
|
||||
_request_body(model_id, prefix, result.request_format, endpoint_tag),
|
||||
headers,
|
||||
timeout,
|
||||
)
|
||||
second = await post(result.request_format)
|
||||
_record(result, second)
|
||||
finally:
|
||||
if owns_client:
|
||||
await client.aclose()
|
||||
await http.aclose()
|
||||
return result
|
||||
|
||||
|
||||
@@ -564,6 +560,7 @@ async def run_cache_checks(
|
||||
timeout: float = PROBE_TIMEOUT_SECONDS,
|
||||
pricing_known: bool = True,
|
||||
endpoint_tag: str | None = None,
|
||||
upstream: "BaseUpstreamProvider | None" = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Run the cache probe and build the three cache/margin rows."""
|
||||
probe = await probe_cache(
|
||||
@@ -573,6 +570,8 @@ async def run_cache_checks(
|
||||
endpoint_tag=endpoint_tag,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
upstream=upstream,
|
||||
model=model,
|
||||
)
|
||||
cost_data = await _price_payload(probe.second_payload, model, provider_fee)
|
||||
return [
|
||||
|
||||
@@ -12,7 +12,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
from routstr.core.db import ModelPathRow
|
||||
from routstr.proxy import reinitialize_upstreams
|
||||
from routstr.upstream.model_paths import encode_model_path
|
||||
from tests.integration.test_certify_endpoint import (
|
||||
|
||||
from .test_certify_endpoint import (
|
||||
_admin_headers,
|
||||
_make_provider,
|
||||
_model_row,
|
||||
|
||||
@@ -22,6 +22,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
from routstr.core.admin import admin_sessions
|
||||
from routstr.core.db import ModelPathRow, ModelRow, UpstreamProviderRow
|
||||
from routstr.proxy import reinitialize_upstreams
|
||||
from routstr.upstream.generic import GenericUpstreamProvider
|
||||
from routstr.upstream.model_paths import encode_model_path
|
||||
|
||||
|
||||
@@ -805,9 +806,7 @@ async def test_certify_explicit_discovered_model_without_override(
|
||||
0.0005,
|
||||
)
|
||||
|
||||
class FakeUpstream:
|
||||
db_id = provider.id
|
||||
|
||||
class FakeUpstream(GenericUpstreamProvider):
|
||||
def get_cached_models(self) -> list[Model]:
|
||||
return [remote_model]
|
||||
|
||||
@@ -820,7 +819,9 @@ async def test_certify_explicit_discovered_model_without_override(
|
||||
return_value=Response(200, json=_mock_chat_response(model="remote-model"))
|
||||
)
|
||||
|
||||
with patch("routstr.proxy.get_upstreams", return_value=[FakeUpstream()]):
|
||||
fake = FakeUpstream(base_url=provider.base_url, api_key=provider.api_key)
|
||||
fake.db_id = provider.id
|
||||
with patch("routstr.proxy.get_upstreams", return_value=[fake]):
|
||||
resp = await integration_client.post(
|
||||
f"/admin/api/upstream-providers/{provider.id}/certify",
|
||||
headers=_admin_headers(),
|
||||
|
||||
@@ -0,0 +1,185 @@
|
||||
"""Certification probes must send the request the proxy would send.
|
||||
|
||||
Each provider type reshapes requests through its hooks (paths, auth headers,
|
||||
query params, model-name transforms). The mocked upstream here answers only
|
||||
the proxy-shaped request, so a probe that hand-builds an OpenAI-style call
|
||||
fails ``endpoint.reachable`` / ``usage.capture`` instead of passing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import respx
|
||||
from httpx import AsyncClient, Response
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
from routstr.proxy import reinitialize_upstreams
|
||||
|
||||
from .test_certify_endpoint import (
|
||||
_admin_headers,
|
||||
_find_row,
|
||||
_mock_chat_response,
|
||||
_model_row,
|
||||
_pin_sats_usd, # noqa: F401 - autouse: pins the sats/USD quote
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Shape:
|
||||
provider_type: str
|
||||
base_url: str
|
||||
model_id: str
|
||||
models_url: str
|
||||
chat_url: str
|
||||
upstream_model: str
|
||||
auth_header: tuple[str, str]
|
||||
params: dict[str, str] = field(default_factory=dict)
|
||||
api_version: str | None = None
|
||||
|
||||
|
||||
SHAPES = [
|
||||
Shape(
|
||||
provider_type="azure",
|
||||
base_url="https://res.openai.azure.com",
|
||||
model_id="gpt-4o",
|
||||
models_url="https://res.openai.azure.com/openai/models",
|
||||
chat_url=(
|
||||
"https://res.openai.azure.com/openai/deployments/gpt-4o/chat/completions"
|
||||
),
|
||||
upstream_model="gpt-4o",
|
||||
auth_header=("api-key", "test-key"),
|
||||
params={"api-version": "2024-10-21"},
|
||||
api_version="2024-10-21",
|
||||
),
|
||||
Shape(
|
||||
provider_type="gemini",
|
||||
base_url="https://generativelanguage.googleapis.com/v1beta",
|
||||
model_id="gemini-2.5-flash",
|
||||
models_url="https://generativelanguage.googleapis.com/v1beta/openai/models",
|
||||
chat_url=(
|
||||
"https://generativelanguage.googleapis.com/v1beta/openai/chat/completions"
|
||||
),
|
||||
upstream_model="gemini-2.5-flash",
|
||||
auth_header=("authorization", "Bearer test-key"),
|
||||
),
|
||||
Shape(
|
||||
provider_type="ollama",
|
||||
base_url="http://ollama.test:11434",
|
||||
model_id="llama3",
|
||||
models_url="http://ollama.test:11434/v1/models",
|
||||
chat_url="http://ollama.test:11434/v1/chat/completions",
|
||||
upstream_model="llama3",
|
||||
auth_header=("authorization", "Bearer test-key"),
|
||||
),
|
||||
Shape(
|
||||
provider_type="anthropic",
|
||||
base_url="https://api.anthropic.com/v1",
|
||||
model_id="claude-sonnet-4.5",
|
||||
models_url="https://api.anthropic.com/v1/models",
|
||||
chat_url="https://api.anthropic.com/v1/chat/completions",
|
||||
upstream_model="claude-sonnet-4-5-20250929",
|
||||
auth_header=("authorization", "Bearer test-key"),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
async def _seed(session: AsyncSession, shape: Shape) -> int:
|
||||
provider = UpstreamProviderRow(
|
||||
provider_type=shape.provider_type,
|
||||
base_url=shape.base_url,
|
||||
api_key="test-key",
|
||||
api_version=shape.api_version,
|
||||
provider_fee=1.0,
|
||||
)
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
await session.refresh(provider)
|
||||
assert provider.id is not None
|
||||
session.add(_model_row(provider.id, model_id=shape.model_id))
|
||||
await session.commit()
|
||||
with patch("routstr.payment.models.sats_usd_price", return_value=0.0005):
|
||||
await reinitialize_upstreams()
|
||||
return provider.id
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("shape", SHAPES, ids=[s.provider_type for s in SHAPES])
|
||||
async def test_certify_matches_proxy_request(
|
||||
shape: Shape,
|
||||
integration_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
with respx.mock(assert_all_called=False) as mock:
|
||||
provider_id = await _seed(integration_session, shape)
|
||||
models_route = mock.get(shape.models_url).mock(
|
||||
return_value=Response(200, json={"data": [{"id": shape.model_id}]})
|
||||
)
|
||||
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(
|
||||
f"/admin/api/upstream-providers/{provider_id}/certify",
|
||||
headers=_admin_headers(),
|
||||
json={"model_id": shape.model_id, "check_cache": True},
|
||||
)
|
||||
|
||||
assert resp.status_code == 200, resp.text
|
||||
rows = resp.json()["rows"]
|
||||
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"
|
||||
|
||||
assert models_route.call_count == 1
|
||||
# The one-token probe plus both cache-probe completions.
|
||||
assert chat_route.call_count == 3
|
||||
header, value = shape.auth_header
|
||||
for call in [*models_route.calls, *chat_route.calls]:
|
||||
assert call.request.headers.get(header) == value
|
||||
for key, expected in shape.params.items():
|
||||
assert call.request.url.params.get(key) == expected
|
||||
for call in chat_route.calls:
|
||||
assert json.loads(call.request.content)["model"] == shape.upstream_model
|
||||
|
||||
|
||||
def test_shape_body_keeps_a_single_cache_control_marker() -> None:
|
||||
"""The cache probe's own marker must not be stamped a second time."""
|
||||
from routstr.payment.models import Architecture, Model, Pricing
|
||||
from routstr.upstream.anthropic import AnthropicUpstreamProvider
|
||||
from routstr.upstream.certification import shape_body
|
||||
from routstr.upstream.certification_cache import _request_body
|
||||
|
||||
model = Model(
|
||||
id="claude-sonnet-4.5",
|
||||
name="claude-sonnet-4.5",
|
||||
created=0,
|
||||
description="",
|
||||
context_length=8192,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="unknown",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(prompt=1e-6, completion=2e-6),
|
||||
sats_pricing=None,
|
||||
per_request_limits=None,
|
||||
top_provider=None,
|
||||
enabled=True,
|
||||
upstream_provider_id=1,
|
||||
canonical_slug=None,
|
||||
)
|
||||
upstream = AnthropicUpstreamProvider(api_key="test-key")
|
||||
body = _request_body(model.id, "prefix", "cache_control", None)
|
||||
|
||||
shaped = shape_body(body, upstream, model)
|
||||
|
||||
assert json.dumps(shaped).count('"cache_control"') == 1
|
||||
assert shaped["model"] == "claude-sonnet-4-5-20250929"
|
||||
@@ -145,7 +145,9 @@ async def test_standalone_preserves_fee_cache_rates_and_usd_fee(
|
||||
return httpx.Response(200, json={"data": [{"id": "test-model"}]})
|
||||
return httpx.Response(200, json={"model": "test-model", "usage": usage})
|
||||
|
||||
kwargs = {"prompt_price": 1e-6, "completion_price": 2e-6} if explicit else {}
|
||||
kwargs: dict[str, Any] = (
|
||||
{"prompt_price": 1e-6, "completion_price": 2e-6} if explicit else {}
|
||||
)
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
|
||||
result = await certify_upstream_url(
|
||||
"https://mock.example/v1",
|
||||
|
||||
@@ -304,6 +304,12 @@ export function ProviderCertificationSetupPanel({
|
||||
Probe prompt caching and margin
|
||||
</Label>
|
||||
</div>
|
||||
{setup.checkCache && (
|
||||
<p className='text-muted-foreground text-xs'>
|
||||
Sends 2–3 extra completions with a ~4.4k-token prompt per model path,
|
||||
billed by the upstream.
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -79,7 +79,6 @@ export type ProviderCertification = z.infer<typeof ProviderCertificationSchema>;
|
||||
export type CertifyProviderRequest = {
|
||||
model_id?: string;
|
||||
model_path?: string;
|
||||
timeout_seconds?: number;
|
||||
check_cache?: boolean;
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user