diff --git a/migrations/versions/e5f6a7b8c9d0_add_model_metadata_to_model_paths.py b/migrations/versions/e5f6a7b8c9d0_add_model_metadata_to_model_paths.py new file mode 100644 index 00000000..5d56c458 --- /dev/null +++ b/migrations/versions/e5f6a7b8c9d0_add_model_metadata_to_model_paths.py @@ -0,0 +1,30 @@ +"""add model metadata to model paths + +Revision ID: e5f6a7b8c9d0 +Revises: b4f7a1c9d2e3 +Create Date: 2026-08-30 00:00:00.000000 +""" + +import sqlalchemy as sa +from alembic import op + +revision = "e5f6a7b8c9d0" +down_revision = "b4f7a1c9d2e3" +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") diff --git a/routstr/core/db.py b/routstr/core/db.py index 4c4d1a62..c9f4268d 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -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", diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 95b6edf2..798df6f4 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -16,6 +16,7 @@ from __future__ import annotations import asyncio import ipaddress +import json import random import time from dataclasses import dataclass @@ -60,10 +61,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) @@ -83,6 +85,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 @@ -263,9 +266,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): @@ -281,6 +289,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: @@ -341,6 +350,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, @@ -348,11 +385,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() @@ -414,6 +450,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): @@ -453,6 +490,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, ) @@ -512,6 +550,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, } @@ -526,6 +565,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, }, ) @@ -740,6 +780,13 @@ def _serialize_path(row: ModelPathRow) -> 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) return { "path": row.path, "provider": { @@ -748,11 +795,12 @@ def _serialize_path(row: ModelPathRow) -> dict[str, Any]: "type": row.provider_type, }, "endpoint": endpoint, + "model": model, } 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( diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 240a4199..32542e89 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -244,6 +244,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, + }, } @@ -367,6 +373,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, @@ -564,9 +601,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 @@ -586,12 +626,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 @@ -765,6 +807,78 @@ 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 +) -> None: + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + endpoint_response = 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", + }, + }, + ], + } + }, + ) + _mock_transport(monkeypatch, lambda request: endpoint_response) + + await mp.refresh_model_paths([provider]) + + payload = await mp.get_all_model_paths() + assert payload["data"][0]["id"] == "claude-opus-4.6" + paths_by_endpoint = { + item["endpoint"]["tag"]: item + for item in payload["data"][0]["paths"] + if item["endpoint"] is not None + } + + anthropic = paths_by_endpoint["anthropic"]["model"] + google = paths_by_endpoint["google-vertex/us"]["model"] + assert anthropic["description"] == "Anthropic's most capable model" + assert google["description"] == "Anthropic's most capable model" + assert anthropic["pricing"] == { + "prompt": "0.000005", + "completion": "0.000025", + } + assert google["pricing"] == { + "prompt": "0.000003", + "completion": "0.000015", + } + assert anthropic["context_length"] == 200_000 + assert google["context_length"] == 128_000 + + @pytest.mark.asyncio async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch