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):
|
class CertifyRequest(BaseModel):
|
||||||
model_id: str | None = None
|
model_id: str | None = None
|
||||||
model_path: str | None = None
|
model_path: str | None = None
|
||||||
timeout_seconds: float | None = None
|
|
||||||
check_cache: bool = True
|
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
|
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
|
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
|
enters the billing path, so it costs nothing from the node's wallet. Its
|
||||||
upstream credit and nothing from the node's wallet.
|
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
|
Returns the read-only report's four ``pricing.*`` rows (re-derived here so
|
||||||
the certification is self-contained), the live rows from
|
the certification is self-contained), the live rows from
|
||||||
@@ -1547,7 +1547,6 @@ async def certify_upstream_provider(
|
|||||||
operator-facing goals.
|
operator-facing goals.
|
||||||
"""
|
"""
|
||||||
from ..upstream.certification import (
|
from ..upstream.certification import (
|
||||||
MAX_PROBE_TIMEOUT_SECONDS,
|
|
||||||
PROBE_TIMEOUT_SECONDS,
|
PROBE_TIMEOUT_SECONDS,
|
||||||
build_checklist,
|
build_checklist,
|
||||||
run_live_checks,
|
run_live_checks,
|
||||||
@@ -1565,6 +1564,7 @@ async def certify_upstream_provider(
|
|||||||
enabled_rows = list(result.all())
|
enabled_rows = list(result.all())
|
||||||
|
|
||||||
endpoint_tag: str | None = None
|
endpoint_tag: str | None = None
|
||||||
|
path_model_id: str | None = None
|
||||||
selected_path: ModelPathRow | None = None
|
selected_path: ModelPathRow | None = None
|
||||||
if payload.model_path is not None:
|
if payload.model_path is not None:
|
||||||
from ..upstream.model_paths import decode_model_path
|
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",
|
detail="Model path is not available for this provider",
|
||||||
)
|
)
|
||||||
endpoint_tag = selector.endpoint_tag
|
endpoint_tag = selector.endpoint_tag
|
||||||
|
path_model_id = selector.model_id
|
||||||
|
|
||||||
evaluations = [
|
evaluations = [
|
||||||
_evaluate_model_row(row, provider, provider_pk) for row in enabled_rows
|
_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
|
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
|
model_obj = None
|
||||||
if model_id:
|
if model_id:
|
||||||
try:
|
try:
|
||||||
@@ -1633,20 +1640,15 @@ async def certify_upstream_provider(
|
|||||||
# The active upstream cache carries the same fee-adjusted USD and sats
|
# 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
|
# pricing used by the proxy, so it is the authoritative fallback for a
|
||||||
# pre-configuration certification probe.
|
# pre-configuration certification probe.
|
||||||
if model_obj is None:
|
if model_obj is None and upstream_obj is not None:
|
||||||
for upstream in get_upstreams():
|
model_obj = next(
|
||||||
if getattr(upstream, "db_id", None) != provider_pk:
|
(
|
||||||
continue
|
model
|
||||||
model_obj = next(
|
for model in upstream_obj.get_cached_models()
|
||||||
(
|
if model.id == model_id or model.forwarded_model_id == model_id
|
||||||
model
|
),
|
||||||
for model in upstream.get_cached_models()
|
None,
|
||||||
if model.id == model_id or model.forwarded_model_id == model_id
|
)
|
||||||
),
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
if model_obj is not None:
|
|
||||||
break
|
|
||||||
if selected_path is not None:
|
if selected_path is not None:
|
||||||
from ..upstream.model_paths import exposed_model_id
|
from ..upstream.model_paths import exposed_model_id
|
||||||
|
|
||||||
@@ -1654,7 +1656,8 @@ async def certify_upstream_provider(
|
|||||||
if (
|
if (
|
||||||
payload.model_id is None
|
payload.model_id is None
|
||||||
or selected_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(
|
raise HTTPException(
|
||||||
status_code=400,
|
status_code=400,
|
||||||
@@ -1726,23 +1729,18 @@ async def certify_upstream_provider(
|
|||||||
provider.provider_fee,
|
provider.provider_fee,
|
||||||
sats_to_usd,
|
sats_to_usd,
|
||||||
)
|
)
|
||||||
# Clamp the admin-supplied timeout per upstream call. The run makes up
|
# The timeout applies per upstream call. The run makes up to five
|
||||||
# to five calls, so the request can stay open for up to five times it.
|
# 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)
|
|
||||||
live_rows = await run_live_checks(
|
live_rows = await run_live_checks(
|
||||||
provider.base_url,
|
provider.base_url,
|
||||||
provider.api_key,
|
provider.api_key,
|
||||||
model_obj,
|
model_obj,
|
||||||
provider_fee=provider.provider_fee,
|
provider_fee=provider.provider_fee,
|
||||||
sats_to_usd=sats_to_usd,
|
sats_to_usd=sats_to_usd,
|
||||||
timeout=timeout,
|
timeout=PROBE_TIMEOUT_SECONDS,
|
||||||
check_cache=payload.check_cache,
|
check_cache=payload.check_cache,
|
||||||
endpoint_tag=endpoint_tag,
|
endpoint_tag=endpoint_tag,
|
||||||
|
upstream=upstream_obj,
|
||||||
)
|
)
|
||||||
|
|
||||||
rows = pricing_rows + live_rows
|
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.
|
ones that must: a ``/models`` heartbeat and a one-token completion.
|
||||||
|
|
||||||
Probes call the upstream directly with ``httpx``, never through the node's
|
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``
|
They sit behind ``POST …/certify`` rather than the read-only ``GET …/report``
|
||||||
because they can block for the length of the timeout.
|
because they can block for the length of the timeout.
|
||||||
"""
|
"""
|
||||||
@@ -25,13 +27,14 @@ from urllib.parse import urlparse
|
|||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from ..core.logging import get_logger
|
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.rates import coerce_rate
|
||||||
from ..payment.usage import normalize_usage
|
from ..payment.usage import normalize_usage
|
||||||
from .model_paths import is_openrouter_base_url
|
from .model_paths import is_openrouter_base_url
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from ..payment.models import Model
|
from ..payment.models import Model
|
||||||
|
from .base import BaseUpstreamProvider
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
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.
|
# Bounded so a dead upstream fails the row rather than wedging the request.
|
||||||
PROBE_TIMEOUT_SECONDS = 15.0
|
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.
|
# The cheapest request that still exercises the usage/cost path.
|
||||||
PROBE_MAX_TOKENS = 1
|
PROBE_MAX_TOKENS = 1
|
||||||
PROBE_PROMPT = "ping"
|
PROBE_PROMPT = "ping"
|
||||||
@@ -181,6 +181,72 @@ class ProbeResult:
|
|||||||
chat_latency_ms: float | None = None
|
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(
|
async def probe_upstream(
|
||||||
base_url: str,
|
base_url: str,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
@@ -189,22 +255,22 @@ async def probe_upstream(
|
|||||||
endpoint_tag: str | None = None,
|
endpoint_tag: str | None = None,
|
||||||
client: httpx.AsyncClient | None = None,
|
client: httpx.AsyncClient | None = None,
|
||||||
timeout: float = PROBE_TIMEOUT_SECONDS,
|
timeout: float = PROBE_TIMEOUT_SECONDS,
|
||||||
|
upstream: "BaseUpstreamProvider | None" = None,
|
||||||
|
model: "Model | None" = None,
|
||||||
) -> ProbeResult:
|
) -> ProbeResult:
|
||||||
"""Call the upstream's ``/models`` and a one-token completion.
|
"""Call the upstream's ``/models`` and a one-token completion.
|
||||||
|
|
||||||
Each HTTP call, including its body read, has an elapsed-time deadline.
|
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 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(
|
result = ProbeResult(
|
||||||
base_url=base_url,
|
base_url=base_url,
|
||||||
models_url=f"{base}/models",
|
models_url=shape.models_url,
|
||||||
chat_url=f"{base}/chat/completions",
|
chat_url=shape.chat_url,
|
||||||
endpoint_tag=endpoint_tag,
|
endpoint_tag=endpoint_tag,
|
||||||
)
|
)
|
||||||
headers = {"Content-Type": "application/json"}
|
headers = shape.headers
|
||||||
if api_key:
|
|
||||||
headers["Authorization"] = f"Bearer {api_key}"
|
|
||||||
|
|
||||||
owns_client = client is None
|
owns_client = client is None
|
||||||
if client is None:
|
if client is None:
|
||||||
@@ -214,7 +280,9 @@ async def probe_upstream(
|
|||||||
started = time.monotonic()
|
started = time.monotonic()
|
||||||
try:
|
try:
|
||||||
async with asyncio.timeout(timeout):
|
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_status = response.status_code
|
||||||
result.models_latency_ms = round((time.monotonic() - started) * 1000, 2)
|
result.models_latency_ms = round((time.monotonic() - started) * 1000, 2)
|
||||||
try:
|
try:
|
||||||
@@ -249,7 +317,10 @@ async def probe_upstream(
|
|||||||
try:
|
try:
|
||||||
async with asyncio.timeout(timeout):
|
async with asyncio.timeout(timeout):
|
||||||
response = await client.post(
|
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_status = response.status_code
|
||||||
result.chat_latency_ms = round((time.monotonic() - started) * 1000, 2)
|
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:
|
def _reported_usd_cost(payload: dict[str, Any]) -> float:
|
||||||
"""The upstream-reported USD cost, or 0.0 when it reported none.
|
"""The upstream-reported USD cost, or 0.0 when it reported none.
|
||||||
|
|
||||||
Mirrors ``_resolve_usd_cost``'s priority and shares ``coerce_rate``, so
|
Uses the engine's own ``_resolve_usd_cost`` so both agree on *which*
|
||||||
this helper and the engine agree on *whether* a cost was reported; only
|
figure is the cost (PPQ.AI BYOK bills ``upstream_inference_cost`` plus the
|
||||||
the arithmetic below is re-derived independently.
|
fee); only the arithmetic below is re-derived independently.
|
||||||
"""
|
"""
|
||||||
usage = payload.get("usage")
|
usage = payload.get("usage")
|
||||||
if not isinstance(usage, dict):
|
if not isinstance(usage, dict):
|
||||||
return 0.0
|
return 0.0
|
||||||
cost_details = usage.get("cost_details")
|
return _resolve_usd_cost(usage, payload)
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
def _fixed_token_pricing_active() -> bool:
|
def _fixed_token_pricing_active() -> bool:
|
||||||
@@ -758,11 +816,14 @@ async def run_live_checks(
|
|||||||
pricing_known: bool = True,
|
pricing_known: bool = True,
|
||||||
check_cache: bool = True,
|
check_cache: bool = True,
|
||||||
endpoint_tag: str | None = None,
|
endpoint_tag: str | None = None,
|
||||||
|
upstream: "BaseUpstreamProvider | None" = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Probe one upstream and build the live/derived rows.
|
"""Probe one upstream and build the live/derived rows.
|
||||||
|
|
||||||
``check_cache`` adds the prompt-cache and margin rows, which cost two
|
``check_cache`` adds the prompt-cache and margin rows, which cost two or
|
||||||
more completions against a long prompt.
|
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(
|
probe = await probe_upstream(
|
||||||
base_url,
|
base_url,
|
||||||
@@ -771,6 +832,8 @@ async def run_live_checks(
|
|||||||
endpoint_tag=endpoint_tag,
|
endpoint_tag=endpoint_tag,
|
||||||
client=client,
|
client=client,
|
||||||
timeout=timeout,
|
timeout=timeout,
|
||||||
|
upstream=upstream,
|
||||||
|
model=model,
|
||||||
)
|
)
|
||||||
rows = [
|
rows = [
|
||||||
safe_row(
|
safe_row(
|
||||||
@@ -863,6 +926,7 @@ async def run_live_checks(
|
|||||||
timeout=timeout,
|
timeout=timeout,
|
||||||
pricing_known=pricing_known,
|
pricing_known=pricing_known,
|
||||||
endpoint_tag=endpoint_tag,
|
endpoint_tag=endpoint_tag,
|
||||||
|
upstream=upstream,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return rows
|
return rows
|
||||||
|
|||||||
@@ -37,11 +37,14 @@ from .certification import (
|
|||||||
_reported_usd_cost,
|
_reported_usd_cost,
|
||||||
_token_rates,
|
_token_rates,
|
||||||
certification_row,
|
certification_row,
|
||||||
|
probe_shape,
|
||||||
safe_row,
|
safe_row,
|
||||||
|
shape_body,
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from ..payment.models import Model
|
from ..payment.models import Model
|
||||||
|
from .base import BaseUpstreamProvider
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
@@ -129,14 +132,15 @@ def _request_body(
|
|||||||
async def _post_completion(
|
async def _post_completion(
|
||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
url: str,
|
url: str,
|
||||||
body: dict[str, Any],
|
body: Any,
|
||||||
headers: dict[str, str],
|
headers: dict[str, str],
|
||||||
|
params: dict[str, str],
|
||||||
timeout: float,
|
timeout: float,
|
||||||
) -> tuple[int | None, dict[str, Any] | None, str | None, float]:
|
) -> tuple[int | None, dict[str, Any] | None, str | None, float]:
|
||||||
started = time.monotonic()
|
started = time.monotonic()
|
||||||
try:
|
try:
|
||||||
async with asyncio.timeout(timeout):
|
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
|
except Exception as exc: # noqa: BLE001 - transport failure is a row status
|
||||||
latency = round((time.monotonic() - started) * 1000, 2)
|
latency = round((time.monotonic() - started) * 1000, 2)
|
||||||
return None, None, f"{type(exc).__name__}: {exc}", latency
|
return None, None, f"{type(exc).__name__}: {exc}", latency
|
||||||
@@ -178,6 +182,8 @@ async def probe_cache(
|
|||||||
endpoint_tag: str | None = None,
|
endpoint_tag: str | None = None,
|
||||||
client: httpx.AsyncClient | None = None,
|
client: httpx.AsyncClient | None = None,
|
||||||
timeout: float = PROBE_TIMEOUT_SECONDS,
|
timeout: float = PROBE_TIMEOUT_SECONDS,
|
||||||
|
upstream: "BaseUpstreamProvider | None" = None,
|
||||||
|
model: "Model | None" = None,
|
||||||
) -> CacheProbeResult:
|
) -> CacheProbeResult:
|
||||||
"""Send the same long prompt twice.
|
"""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
|
retry on HTTP 400/422, and the second call mirrors whichever format
|
||||||
succeeded. Each call's elapsed deadline includes the response body.
|
succeeded. Each call's elapsed deadline includes the response body.
|
||||||
"""
|
"""
|
||||||
base = base_url.rstrip("/")
|
shape = probe_shape(base_url, api_key, upstream, model)
|
||||||
result = CacheProbeResult(
|
result = CacheProbeResult(chat_url=shape.chat_url, endpoint_tag=endpoint_tag)
|
||||||
chat_url=f"{base}/chat/completions", endpoint_tag=endpoint_tag
|
|
||||||
)
|
|
||||||
headers = {"Content-Type": "application/json"}
|
|
||||||
if api_key:
|
|
||||||
headers["Authorization"] = f"Bearer {api_key}"
|
|
||||||
prefix = cache_probe_prefix()
|
prefix = cache_probe_prefix()
|
||||||
|
|
||||||
owns_client = client is None
|
owns_client = client is None
|
||||||
if client is None:
|
http = client if client is not None else httpx.AsyncClient(timeout=timeout)
|
||||||
client = httpx.AsyncClient(timeout=timeout)
|
|
||||||
try:
|
async def post(
|
||||||
first = await _post_completion(
|
fmt: str,
|
||||||
client,
|
) -> 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,
|
result.chat_url,
|
||||||
_request_body(model_id, prefix, "cache_control", endpoint_tag),
|
shape_body(body, upstream, model),
|
||||||
headers,
|
shape.headers,
|
||||||
|
shape.chat_params,
|
||||||
timeout,
|
timeout,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
first = await post("cache_control")
|
||||||
if first[0] in (400, 422):
|
if first[0] in (400, 422):
|
||||||
result.request_format = "plain"
|
result.request_format = "plain"
|
||||||
first = await _post_completion(
|
first = await post("plain")
|
||||||
client,
|
|
||||||
result.chat_url,
|
|
||||||
_request_body(model_id, prefix, "plain", endpoint_tag),
|
|
||||||
headers,
|
|
||||||
timeout,
|
|
||||||
)
|
|
||||||
_record(result, first)
|
_record(result, first)
|
||||||
if not _is_2xx(first[0]):
|
if not _is_2xx(first[0]):
|
||||||
return result
|
return result
|
||||||
second = await _post_completion(
|
second = await post(result.request_format)
|
||||||
client,
|
|
||||||
result.chat_url,
|
|
||||||
_request_body(model_id, prefix, result.request_format, endpoint_tag),
|
|
||||||
headers,
|
|
||||||
timeout,
|
|
||||||
)
|
|
||||||
_record(result, second)
|
_record(result, second)
|
||||||
finally:
|
finally:
|
||||||
if owns_client:
|
if owns_client:
|
||||||
await client.aclose()
|
await http.aclose()
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
@@ -564,6 +560,7 @@ async def run_cache_checks(
|
|||||||
timeout: float = PROBE_TIMEOUT_SECONDS,
|
timeout: float = PROBE_TIMEOUT_SECONDS,
|
||||||
pricing_known: bool = True,
|
pricing_known: bool = True,
|
||||||
endpoint_tag: str | None = None,
|
endpoint_tag: str | None = None,
|
||||||
|
upstream: "BaseUpstreamProvider | None" = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Run the cache probe and build the three cache/margin rows."""
|
"""Run the cache probe and build the three cache/margin rows."""
|
||||||
probe = await probe_cache(
|
probe = await probe_cache(
|
||||||
@@ -573,6 +570,8 @@ async def run_cache_checks(
|
|||||||
endpoint_tag=endpoint_tag,
|
endpoint_tag=endpoint_tag,
|
||||||
client=client,
|
client=client,
|
||||||
timeout=timeout,
|
timeout=timeout,
|
||||||
|
upstream=upstream,
|
||||||
|
model=model,
|
||||||
)
|
)
|
||||||
cost_data = await _price_payload(probe.second_payload, model, provider_fee)
|
cost_data = await _price_payload(probe.second_payload, model, provider_fee)
|
||||||
return [
|
return [
|
||||||
|
|||||||
@@ -12,7 +12,8 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
|||||||
from routstr.core.db import ModelPathRow
|
from routstr.core.db import ModelPathRow
|
||||||
from routstr.proxy import reinitialize_upstreams
|
from routstr.proxy import reinitialize_upstreams
|
||||||
from routstr.upstream.model_paths import encode_model_path
|
from routstr.upstream.model_paths import encode_model_path
|
||||||
from tests.integration.test_certify_endpoint import (
|
|
||||||
|
from .test_certify_endpoint import (
|
||||||
_admin_headers,
|
_admin_headers,
|
||||||
_make_provider,
|
_make_provider,
|
||||||
_model_row,
|
_model_row,
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession
|
|||||||
from routstr.core.admin import admin_sessions
|
from routstr.core.admin import admin_sessions
|
||||||
from routstr.core.db import ModelPathRow, ModelRow, UpstreamProviderRow
|
from routstr.core.db import ModelPathRow, ModelRow, UpstreamProviderRow
|
||||||
from routstr.proxy import reinitialize_upstreams
|
from routstr.proxy import reinitialize_upstreams
|
||||||
|
from routstr.upstream.generic import GenericUpstreamProvider
|
||||||
from routstr.upstream.model_paths import encode_model_path
|
from routstr.upstream.model_paths import encode_model_path
|
||||||
|
|
||||||
|
|
||||||
@@ -805,9 +806,7 @@ async def test_certify_explicit_discovered_model_without_override(
|
|||||||
0.0005,
|
0.0005,
|
||||||
)
|
)
|
||||||
|
|
||||||
class FakeUpstream:
|
class FakeUpstream(GenericUpstreamProvider):
|
||||||
db_id = provider.id
|
|
||||||
|
|
||||||
def get_cached_models(self) -> list[Model]:
|
def get_cached_models(self) -> list[Model]:
|
||||||
return [remote_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"))
|
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(
|
resp = await integration_client.post(
|
||||||
f"/admin/api/upstream-providers/{provider.id}/certify",
|
f"/admin/api/upstream-providers/{provider.id}/certify",
|
||||||
headers=_admin_headers(),
|
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={"data": [{"id": "test-model"}]})
|
||||||
return httpx.Response(200, json={"model": "test-model", "usage": usage})
|
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:
|
async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as client:
|
||||||
result = await certify_upstream_url(
|
result = await certify_upstream_url(
|
||||||
"https://mock.example/v1",
|
"https://mock.example/v1",
|
||||||
|
|||||||
@@ -304,6 +304,12 @@ export function ProviderCertificationSetupPanel({
|
|||||||
Probe prompt caching and margin
|
Probe prompt caching and margin
|
||||||
</Label>
|
</Label>
|
||||||
</div>
|
</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>
|
</div>
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -79,7 +79,6 @@ export type ProviderCertification = z.infer<typeof ProviderCertificationSchema>;
|
|||||||
export type CertifyProviderRequest = {
|
export type CertifyProviderRequest = {
|
||||||
model_id?: string;
|
model_id?: string;
|
||||||
model_path?: string;
|
model_path?: string;
|
||||||
timeout_seconds?: number;
|
|
||||||
check_cache?: boolean;
|
check_cache?: boolean;
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user