mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Merge pull request #719 from Routstr/add-model-path-metadata
add model path metadata
This commit is contained in:
@@ -0,0 +1,27 @@
|
||||
"""add model_metadata to model_paths
|
||||
|
||||
Revision ID: a3f1b6c204de
|
||||
Revises: e5a6b7c8d9f0
|
||||
Create Date: 2026-09-15 21:50:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "a3f1b6c204de"
|
||||
down_revision = "e5a6b7c8d9f0"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"model_paths",
|
||||
sa.Column("model_metadata", sa.Text(), nullable=False, server_default="{}"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("model_paths", "model_metadata")
|
||||
+8
-3
@@ -431,9 +431,10 @@ class ModelRow(SQLModel, table=True): # type: ignore
|
||||
class ModelPathRow(SQLModel, table=True): # type: ignore
|
||||
"""Upstream provider path a model is reachable through.
|
||||
|
||||
Discovery/visibility data only. ``model_id`` is intentionally NOT globally
|
||||
unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or
|
||||
id``) grouped across every provider that exposes the model. A single model
|
||||
Discovery data plus provider-specific model metadata. ``model_id`` is
|
||||
intentionally NOT globally unique: it is the client-visible ``/v1/models``
|
||||
id (``forwarded_model_id or id``) grouped across every provider that exposes
|
||||
the model. A single model
|
||||
can therefore have several rows — one per direct provider path plus one per
|
||||
OpenRouter sub-provider endpoint.
|
||||
"""
|
||||
@@ -470,6 +471,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore
|
||||
endpoint_name: str | None = Field(
|
||||
default=None, description="Human-readable endpoint display name"
|
||||
)
|
||||
model_metadata: str = Field(
|
||||
default="{}",
|
||||
description="JSON model metadata specific to this provider path",
|
||||
)
|
||||
upstream_provider_id: int = Field(
|
||||
index=True,
|
||||
foreign_key="upstream_providers.id",
|
||||
|
||||
+134
-11
@@ -13,6 +13,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
@@ -57,10 +58,11 @@ ModelKey = tuple[str, int]
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EndpointIdentity:
|
||||
"""Exact OpenRouter endpoint identity returned by ``/endpoints``."""
|
||||
"""Exact OpenRouter endpoint and its provider-specific model metadata."""
|
||||
|
||||
tag: str
|
||||
provider_name: str | None
|
||||
model_metadata: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -80,6 +82,7 @@ class DiscoveredPath:
|
||||
model_id: str
|
||||
path: str
|
||||
provider: ConfiguredProviderIdentity
|
||||
model_metadata: dict[str, Any]
|
||||
endpoint_tag: str | None = None
|
||||
endpoint_name: str | None = None
|
||||
|
||||
@@ -307,9 +310,14 @@ async def _fetch_openrouter_endpoint_subproviders(
|
||||
try:
|
||||
payload = resp.json()
|
||||
data = payload.get("data") if isinstance(payload, dict) else None
|
||||
endpoints = data.get("endpoints") if isinstance(data, dict) else None
|
||||
if not isinstance(data, dict):
|
||||
raise ValueError("data must be an object")
|
||||
endpoints = data.get("endpoints")
|
||||
if not isinstance(endpoints, list):
|
||||
raise ValueError("endpoints must be a list")
|
||||
common_metadata = {
|
||||
key: value for key, value in data.items() if key != "endpoints"
|
||||
}
|
||||
identities: dict[str, EndpointIdentity] = {}
|
||||
for endpoint in endpoints:
|
||||
if not isinstance(endpoint, dict):
|
||||
@@ -325,6 +333,7 @@ async def _fetch_openrouter_endpoint_subproviders(
|
||||
provider_name=provider_name
|
||||
if isinstance(provider_name, str) and provider_name
|
||||
else None,
|
||||
model_metadata={**common_metadata, **endpoint},
|
||||
),
|
||||
)
|
||||
if endpoints and not identities:
|
||||
@@ -385,6 +394,34 @@ async def _load_model_visibility() -> tuple[
|
||||
return overrides_by_key, disabled_model_keys, provider_identities
|
||||
|
||||
|
||||
def _serialize_model_metadata(model: object, model_id: str) -> dict[str, Any]:
|
||||
"""Serialize provider-specific model details into the public API shape."""
|
||||
model_dict = getattr(model, "dict", None)
|
||||
if callable(model_dict):
|
||||
metadata = dict(model_dict())
|
||||
else:
|
||||
metadata = {
|
||||
key: value for key, value in vars(model).items() if not key.startswith("_")
|
||||
}
|
||||
|
||||
for field in (
|
||||
"architecture",
|
||||
"pricing",
|
||||
"sats_pricing",
|
||||
"per_request_limits",
|
||||
"top_provider",
|
||||
"alias_ids",
|
||||
):
|
||||
value = metadata.get(field)
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
metadata[field] = json.loads(value)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
metadata["id"] = model_id
|
||||
return metadata
|
||||
|
||||
|
||||
def _apply_model_visibility(
|
||||
upstream: BaseUpstreamProvider,
|
||||
overrides_by_key: dict[ModelKey, ModelRow] | None,
|
||||
@@ -392,11 +429,10 @@ def _apply_model_visibility(
|
||||
) -> list[object]:
|
||||
"""Return provider models after DB disabled/override state is applied.
|
||||
|
||||
Only the identity fields (``id``, ``forwarded_model_id``,
|
||||
``canonical_slug``) matter for path discovery, so DB override rows are used
|
||||
directly rather than rebuilt into fully priced ``Model`` objects — the
|
||||
pricing pipeline costs ~0.7ms of event-loop CPU per row for data this
|
||||
module immediately discards.
|
||||
DB override rows are used directly rather than rebuilt into priced
|
||||
``Model`` objects. Their JSON metadata fields are decoded when each path is
|
||||
collected, preserving the provider-specific stored values without running
|
||||
the routing price-selection pipeline.
|
||||
"""
|
||||
overrides_by_key = overrides_by_key or {}
|
||||
disabled_model_keys = disabled_model_keys or set()
|
||||
@@ -458,6 +494,7 @@ async def _collect_provider_paths(
|
||||
provider_identity.base_url, provider_identity.id, model_id
|
||||
),
|
||||
provider=provider_identity,
|
||||
model_metadata=_serialize_model_metadata(model, model_id),
|
||||
)
|
||||
|
||||
if not is_openrouter_base_url(upstream.base_url):
|
||||
@@ -497,6 +534,7 @@ async def _collect_provider_paths(
|
||||
endpoint.tag,
|
||||
),
|
||||
provider=provider_identity,
|
||||
model_metadata={**endpoint.model_metadata, "id": model_id},
|
||||
endpoint_tag=endpoint.tag,
|
||||
endpoint_name=endpoint.provider_name,
|
||||
)
|
||||
@@ -556,6 +594,7 @@ async def _persist_provider_paths(
|
||||
"provider_type": discovered.provider.provider_type,
|
||||
"endpoint_tag": discovered.endpoint_tag,
|
||||
"endpoint_name": discovered.endpoint_name,
|
||||
"model_metadata": json.dumps(discovered.model_metadata),
|
||||
"upstream_provider_id": upstream_provider_id,
|
||||
"updated_at": now,
|
||||
}
|
||||
@@ -570,6 +609,7 @@ async def _persist_provider_paths(
|
||||
"provider_type": insert_stmt.excluded.provider_type,
|
||||
"endpoint_tag": insert_stmt.excluded.endpoint_tag,
|
||||
"endpoint_name": insert_stmt.excluded.endpoint_name,
|
||||
"model_metadata": insert_stmt.excluded.model_metadata,
|
||||
"updated_at": insert_stmt.excluded.updated_at,
|
||||
},
|
||||
)
|
||||
@@ -780,10 +820,83 @@ async def refresh_model_paths_periodically(
|
||||
break
|
||||
|
||||
|
||||
def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
|
||||
def _price_in_sats(model: dict[str, Any], provider_fee: float) -> None:
|
||||
"""Run a path's USD rates through the ``/v1/models`` pricing pipeline.
|
||||
|
||||
Metadata copied from the provider model cache is already priced. OpenRouter
|
||||
endpoint metadata is not: it carries that endpoint's own USD rates, which
|
||||
still need the cache backfill, the provider fee and the sats conversion.
|
||||
"""
|
||||
pricing = model.get("pricing")
|
||||
if model.get("sats_pricing") or not isinstance(pricing, dict):
|
||||
return
|
||||
|
||||
from ..payment.models import (
|
||||
Architecture,
|
||||
Model,
|
||||
Pricing,
|
||||
TopProvider,
|
||||
_calculate_usd_max_costs,
|
||||
_update_model_sats_pricing,
|
||||
backfill_cache_pricing,
|
||||
)
|
||||
from ..payment.price import sats_usd_price
|
||||
|
||||
try:
|
||||
model_id = model.get("forwarded_model_id") or model["id"]
|
||||
usd = backfill_cache_pricing(model_id, Pricing.parse_obj(pricing))
|
||||
usd = Pricing.parse_obj({k: v * provider_fee for k, v in usd.dict().items()})
|
||||
priced = Model(
|
||||
id=model_id,
|
||||
name=model.get("name") or model_id,
|
||||
created=0,
|
||||
description="",
|
||||
context_length=model.get("context_length") or 0,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=[],
|
||||
output_modalities=[],
|
||||
tokenizer="",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=usd,
|
||||
top_provider=TopProvider(
|
||||
context_length=model.get("context_length"),
|
||||
max_completion_tokens=model.get("max_completion_tokens"),
|
||||
),
|
||||
)
|
||||
(
|
||||
usd.max_prompt_cost,
|
||||
usd.max_completion_cost,
|
||||
usd.max_cost,
|
||||
) = _calculate_usd_max_costs(priced)
|
||||
priced = _update_model_sats_pricing(priced, sats_usd_price())
|
||||
except Exception as exc:
|
||||
# An endpoint with rates we cannot price is still a usable route, so it
|
||||
# is served with its raw upstream pricing rather than dropped.
|
||||
logger.warning(
|
||||
"Could not calculate sats pricing for model path",
|
||||
extra={"model_id": model.get("id"), "error": str(exc)},
|
||||
)
|
||||
return
|
||||
|
||||
if priced.sats_pricing:
|
||||
model["pricing"] = usd.dict()
|
||||
model["sats_pricing"] = priced.sats_pricing.dict()
|
||||
|
||||
|
||||
def _serialize_path(row: ModelPathRow, provider_fee: float) -> dict[str, Any]:
|
||||
endpoint = None
|
||||
if row.endpoint_tag or row.endpoint_name:
|
||||
endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name}
|
||||
try:
|
||||
model = json.loads(row.model_metadata)
|
||||
except (TypeError, ValueError):
|
||||
model = {}
|
||||
if not isinstance(model, dict):
|
||||
model = {}
|
||||
model.setdefault("id", row.model_id)
|
||||
_price_in_sats(model, provider_fee)
|
||||
return {
|
||||
"path": row.path,
|
||||
"provider": {
|
||||
@@ -792,11 +905,17 @@ def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
|
||||
"type": row.provider_type,
|
||||
},
|
||||
"endpoint": endpoint,
|
||||
"model": model,
|
||||
}
|
||||
|
||||
|
||||
async def _provider_fees(session: "AsyncSession") -> dict[int, float]:
|
||||
rows = (await session.exec(select(UpstreamProviderRow))).all()
|
||||
return {row.id: row.provider_fee for row in rows if row.id is not None}
|
||||
|
||||
|
||||
async def get_all_model_paths() -> dict:
|
||||
"""All models with their exact selectable routes."""
|
||||
"""All models with exact routes and provider-specific model metadata."""
|
||||
async with create_session() as session:
|
||||
rows = (
|
||||
await session.exec(
|
||||
@@ -807,6 +926,7 @@ async def get_all_model_paths() -> dict:
|
||||
)
|
||||
)
|
||||
).all()
|
||||
fees = await _provider_fees(session)
|
||||
|
||||
grouped: dict[str, list[dict[str, Any]]] = {}
|
||||
seen_paths: dict[str, set[str]] = {}
|
||||
@@ -816,7 +936,9 @@ async def get_all_model_paths() -> dict:
|
||||
if row.path in seen_paths.setdefault(row.model_id, set()):
|
||||
continue
|
||||
seen_paths[row.model_id].add(row.path)
|
||||
grouped.setdefault(row.model_id, []).append(_serialize_path(row))
|
||||
grouped.setdefault(row.model_id, []).append(
|
||||
_serialize_path(row, fees.get(row.upstream_provider_id, 1.01))
|
||||
)
|
||||
data = [
|
||||
{
|
||||
"id": grouped_model_id,
|
||||
@@ -850,6 +972,7 @@ async def get_paths_for_model(model_id: str) -> dict:
|
||||
unprefixed_id = public_model_id(model_id)
|
||||
if unprefixed_id != model_id:
|
||||
rows = await load_rows(session, unprefixed_id)
|
||||
fees = await _provider_fees(session)
|
||||
|
||||
seen: set[str] = set()
|
||||
paths: list[dict] = []
|
||||
@@ -859,5 +982,5 @@ async def get_paths_for_model(model_id: str) -> dict:
|
||||
if row.path in seen:
|
||||
continue
|
||||
seen.add(row.path)
|
||||
paths.append(_serialize_path(row))
|
||||
paths.append(_serialize_path(row, fees.get(row.upstream_provider_id, 1.01)))
|
||||
return {"data": paths, "updated_at": updated_at or None}
|
||||
|
||||
@@ -28,6 +28,7 @@ os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
|
||||
from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402
|
||||
from routstr.payment import price as price_module # noqa: E402
|
||||
from routstr.payment.models import models_router # noqa: E402
|
||||
from routstr.upstream import model_paths as mp # noqa: E402
|
||||
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
|
||||
@@ -202,6 +203,75 @@ async def patched_session(
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
# 1 sat = $0.00005, the quote the path pricing converts with in these tests.
|
||||
_QUOTE = 5.0e-5
|
||||
_DEFAULT_FEE = 1.01
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sats_quote(monkeypatch: pytest.MonkeyPatch) -> float:
|
||||
monkeypatch.setattr(price_module, "sats_usd_price", lambda: _QUOTE)
|
||||
return _QUOTE
|
||||
|
||||
|
||||
def _priced_endpoints_response() -> httpx.Response:
|
||||
"""Two endpoints for one model, priced and sized differently."""
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"data": {
|
||||
"id": "anthropic/claude-opus-4.6",
|
||||
"name": "Claude Opus 4.6",
|
||||
"description": "Anthropic's most capable model",
|
||||
"architecture": {
|
||||
"input_modalities": ["text", "image"],
|
||||
"output_modalities": ["text"],
|
||||
"tokenizer": "Claude",
|
||||
"instruct_type": None,
|
||||
},
|
||||
"endpoints": [
|
||||
{
|
||||
"provider_name": "Anthropic",
|
||||
"tag": "anthropic",
|
||||
"context_length": 200_000,
|
||||
"pricing": {
|
||||
"prompt": "0.000005",
|
||||
"completion": "0.000025",
|
||||
},
|
||||
},
|
||||
{
|
||||
"provider_name": "Google",
|
||||
"tag": "google-vertex/us",
|
||||
"context_length": 128_000,
|
||||
"pricing": {
|
||||
"prompt": "0.000003",
|
||||
"completion": "0.000015",
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _models_by_endpoint(payload: dict, model_id: str) -> dict[str, dict]:
|
||||
entry = next(item for item in payload["data"] if item["id"] == model_id)
|
||||
return {
|
||||
path["endpoint"]["tag"]: path["model"]
|
||||
for path in entry["paths"]
|
||||
if path["endpoint"] is not None
|
||||
}
|
||||
|
||||
|
||||
async def _set_provider_fee(engine: AsyncEngine, provider_id: int, fee: float) -> None:
|
||||
async with AsyncSession(engine) as session:
|
||||
provider = await session.get(UpstreamProviderRow, provider_id)
|
||||
assert provider is not None
|
||||
provider.provider_fee = fee
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
|
||||
|
||||
def _paths_of(payload: dict, model_id: str) -> set[str]:
|
||||
for entry in payload["data"]:
|
||||
if entry["id"] == model_id:
|
||||
@@ -244,6 +314,12 @@ def _path_entry(
|
||||
or ("anthropic" if provider_id == 1 else "openrouter"),
|
||||
},
|
||||
"endpoint": endpoint,
|
||||
"model": {
|
||||
"id": model_id,
|
||||
"forwarded_model_id": None,
|
||||
"canonical_slug": None,
|
||||
"enabled": True,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@@ -379,6 +455,37 @@ async def test_direct_provider_single_path_uses_provider_type(
|
||||
assert payload["updated_at"] is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_all_model_paths_includes_details_for_each_path(
|
||||
patched_session: AsyncEngine,
|
||||
) -> None:
|
||||
model = _model("claude-opus-4.6")
|
||||
model.name = "Claude Opus 4.6"
|
||||
model.description = "Anthropic's most capable model"
|
||||
model.pricing = {"prompt": 0.000001, "completion": 0.000002}
|
||||
provider = _FakeProvider(
|
||||
provider_type="anthropic",
|
||||
base_url="https://api.anthropic.com/v1",
|
||||
models=[model],
|
||||
db_id=1,
|
||||
)
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
payload = await mp.get_all_model_paths()
|
||||
|
||||
expected_path = _path_entry(1, "claude-opus-4.6")
|
||||
expected_path["model"] = {
|
||||
"id": "claude-opus-4.6",
|
||||
"forwarded_model_id": None,
|
||||
"canonical_slug": None,
|
||||
"enabled": True,
|
||||
"name": "Claude Opus 4.6",
|
||||
"description": "Anthropic's most capable model",
|
||||
"pricing": {"prompt": 0.000001, "completion": 0.000002},
|
||||
}
|
||||
assert payload["data"] == [{"id": "claude-opus-4.6", "paths": [expected_path]}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_direct_path_masks_private_configured_provider_url(
|
||||
patched_session: AsyncEngine,
|
||||
@@ -576,9 +683,12 @@ async def test_refresh_model_paths_uses_db_forwarded_alias(
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
assert (await mp.get_all_model_paths())["data"] == [
|
||||
{"id": "public-alias", "paths": [_path_entry(1, "public-alias")]}
|
||||
]
|
||||
payload = await mp.get_all_model_paths()
|
||||
assert _paths_of(payload, "public-alias") == {_expected_path(1, "public-alias")}
|
||||
model = payload["data"][0]["paths"][0]["model"]
|
||||
assert model["id"] == "public-alias"
|
||||
assert model["description"] == "test model"
|
||||
assert model["pricing"] == {"prompt": 0.000001, "completion": 0.000002}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -598,12 +708,14 @@ async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cac
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
assert (await mp.get_all_model_paths())["data"] == [
|
||||
{
|
||||
"id": "public-deployment",
|
||||
"paths": [_path_entry(1, "public-deployment")],
|
||||
}
|
||||
]
|
||||
payload = await mp.get_all_model_paths()
|
||||
assert _paths_of(payload, "public-deployment") == {
|
||||
_expected_path(1, "public-deployment")
|
||||
}
|
||||
model = payload["data"][0]["paths"][0]["model"]
|
||||
assert model["id"] == "public-deployment"
|
||||
assert model["description"] == "test model"
|
||||
assert model["context_length"] == 8192
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -777,6 +889,148 @@ async def test_openrouter_provider_adds_endpoint_paths(
|
||||
assert {item["provider"]["id"] for item in payload["data"]} == {2}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openrouter_paths_include_endpoint_specific_model_prices(
|
||||
patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float
|
||||
) -> None:
|
||||
provider = _FakeOpenRouterProvider(
|
||||
models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")],
|
||||
db_id=2,
|
||||
)
|
||||
_mock_transport(monkeypatch, lambda request: _priced_endpoints_response())
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
payload = await mp.get_all_model_paths()
|
||||
assert payload["data"][0]["id"] == "claude-opus-4.6"
|
||||
models = _models_by_endpoint(payload, "claude-opus-4.6")
|
||||
|
||||
anthropic = models["anthropic"]
|
||||
google = models["google-vertex/us"]
|
||||
assert anthropic["description"] == "Anthropic's most capable model"
|
||||
assert google["description"] == "Anthropic's most capable model"
|
||||
assert anthropic["pricing"]["prompt"] == pytest.approx(0.000005 * _DEFAULT_FEE)
|
||||
assert anthropic["pricing"]["completion"] == pytest.approx(0.000025 * _DEFAULT_FEE)
|
||||
assert google["pricing"]["prompt"] == pytest.approx(0.000003 * _DEFAULT_FEE)
|
||||
assert google["pricing"]["completion"] == pytest.approx(0.000015 * _DEFAULT_FEE)
|
||||
assert anthropic["context_length"] == 200_000
|
||||
assert google["context_length"] == 128_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_endpoint_paths_are_priced_in_sats_from_their_own_rates(
|
||||
patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float
|
||||
) -> None:
|
||||
provider = _FakeOpenRouterProvider(
|
||||
models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")],
|
||||
db_id=2,
|
||||
)
|
||||
_mock_transport(monkeypatch, lambda request: _priced_endpoints_response())
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6")
|
||||
anthropic = models["anthropic"]["sats_pricing"]
|
||||
google = models["google-vertex/us"]["sats_pricing"]
|
||||
|
||||
assert anthropic["prompt"] == pytest.approx(0.000005 * _DEFAULT_FEE / sats_quote)
|
||||
assert anthropic["completion"] == pytest.approx(
|
||||
0.000025 * _DEFAULT_FEE / sats_quote
|
||||
)
|
||||
assert google["prompt"] == pytest.approx(0.000003 * _DEFAULT_FEE / sats_quote)
|
||||
assert google["completion"] == pytest.approx(0.000015 * _DEFAULT_FEE / sats_quote)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_endpoint_max_costs_use_that_endpoint_context_length(
|
||||
patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float
|
||||
) -> None:
|
||||
provider = _FakeOpenRouterProvider(
|
||||
models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")],
|
||||
db_id=2,
|
||||
)
|
||||
_mock_transport(monkeypatch, lambda request: _priced_endpoints_response())
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6")
|
||||
# Max cost is the context window billed at the dearer of the two rates.
|
||||
assert models["anthropic"]["sats_pricing"]["max_cost"] == pytest.approx(
|
||||
200_000 * 0.000025 * _DEFAULT_FEE / sats_quote
|
||||
)
|
||||
assert models["google-vertex/us"]["sats_pricing"]["max_cost"] == pytest.approx(
|
||||
128_000 * 0.000015 * _DEFAULT_FEE / sats_quote
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_path_pricing_uses_the_provider_fee_of_its_own_provider(
|
||||
patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch, sats_quote: float
|
||||
) -> None:
|
||||
await _set_provider_fee(patched_session, 2, 1.5)
|
||||
provider = _FakeOpenRouterProvider(
|
||||
models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")],
|
||||
db_id=2,
|
||||
)
|
||||
_mock_transport(monkeypatch, lambda request: _priced_endpoints_response())
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6")
|
||||
anthropic = models["anthropic"]
|
||||
assert anthropic["pricing"]["prompt"] == pytest.approx(0.000005 * 1.5)
|
||||
assert anthropic["sats_pricing"]["prompt"] == pytest.approx(
|
||||
0.000005 * 1.5 / sats_quote
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paths_keep_upstream_pricing_when_the_quote_is_unavailable(
|
||||
patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def _no_quote() -> float:
|
||||
raise ValueError("SATS price not initialized")
|
||||
|
||||
monkeypatch.setattr(price_module, "sats_usd_price", _no_quote)
|
||||
provider = _FakeOpenRouterProvider(
|
||||
models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")],
|
||||
db_id=2,
|
||||
)
|
||||
_mock_transport(monkeypatch, lambda request: _priced_endpoints_response())
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
models = _models_by_endpoint(await mp.get_all_model_paths(), "claude-opus-4.6")
|
||||
anthropic = models["anthropic"]
|
||||
assert "sats_pricing" not in anthropic
|
||||
assert anthropic["pricing"] == {
|
||||
"prompt": "0.000005",
|
||||
"completion": "0.000025",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_already_priced_metadata_is_not_priced_again(
|
||||
patched_session: AsyncEngine, sats_quote: float
|
||||
) -> None:
|
||||
model = _model("claude-opus-4.6")
|
||||
model.pricing = {"prompt": 0.000001, "completion": 0.000002}
|
||||
model.sats_pricing = {"prompt": 0.02, "completion": 0.04}
|
||||
provider = _FakeProvider(
|
||||
provider_type="anthropic",
|
||||
base_url="https://api.anthropic.com/v1",
|
||||
models=[model],
|
||||
db_id=1,
|
||||
)
|
||||
|
||||
await mp.refresh_model_paths([provider])
|
||||
|
||||
payload = await mp.get_all_model_paths()
|
||||
priced = payload["data"][0]["paths"][0]["model"]
|
||||
assert priced["sats_pricing"] == {"prompt": 0.02, "completion": 0.04}
|
||||
assert priced["pricing"] == {"prompt": 0.000001, "completion": 0.000002}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openrouter_uses_exact_tag_even_when_display_name_is_router(
|
||||
patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch
|
||||
|
||||
Reference in New Issue
Block a user