Merge pull request #719 from Routstr/add-model-path-metadata

add model path metadata
This commit is contained in:
9qeklajc
2026-09-15 22:00:29 +02:00
committed by GitHub
4 changed files with 432 additions and 23 deletions
@@ -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
View File
@@ -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
View File
@@ -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}
+263 -9
View File
@@ -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