From 349d8dd009750798b50eda5675c9cb175e7fc673 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 5 Jul 2026 23:55:49 +0200 Subject: [PATCH 01/14] add model path endpoint --- docs/api/endpoints.md | 52 ++ docs/api/overview.md | 3 +- docs/provider/configuration.md | 6 + .../d7e8f9a0b1c2_add_model_paths_table.py | 66 ++ routstr/core/db.py | 35 ++ routstr/core/main.py | 11 + routstr/core/settings.py | 3 + routstr/payment/models.py | 23 + routstr/upstream/model_paths.py | 348 +++++++++++ tests/unit/test_model_paths.py | 591 ++++++++++++++++++ 10 files changed, 1137 insertions(+), 1 deletion(-) create mode 100644 migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py create mode 100644 routstr/upstream/model_paths.py create mode 100644 tests/unit/test_model_paths.py diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 5a9ea68d..efe57a68 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -327,6 +327,58 @@ GET /v1/models } ``` +### List Model Paths + +Get the upstream provider paths each advertised model can be reached through. +This is discovery data only; routing still chooses the provider per request. + +```http +GET /v1/models/paths +``` + +**Response:** + +```json +{ + "data": [ + { + "id": "claude-sonnet-4", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"} + ] + } + ] +} +``` + +### List Paths for One Model + +Use a query parameter so model IDs containing `/` are handled safely. Lookup is +by the public, unqualified model ID: `glm-5v-turbo` resolves +`z-ai/glm-5v-turbo`, and `deepseek-v4-pro` and `deepseek/deepseek-v4-pro` +return the same merged path set. + +```http +GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 +``` + +**Response:** + +```json +{ + "data": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"} + ] +} +``` + +Model IDs in responses are unqualified display IDs: provider prefixes such as +`z-ai/` or `openai/` are stripped. Path values match the provider string stamped +on chat-completion responses, such as `anthropic`, `generic:my-upstream`, or +`openrouter:Anthropic`. + ## Wallet Management ### Create Wallet (Coming Soon) diff --git a/docs/api/overview.md b/docs/api/overview.md index 92fedd7e..c82e4e22 100644 --- a/docs/api/overview.md +++ b/docs/api/overview.md @@ -100,6 +100,7 @@ All errors follow a consistent format: Standard OpenAI-compatible endpoints: - **Models**: `/v1/models` +- **Model paths**: `/v1/models/paths`, `/v1/models/paths/model?model_id=...` - **Responses**: `/v1/responses` - **Chat Completions**: `/v1/chat/completions` - **Embeddings**: `/v1/embeddings` @@ -302,7 +303,7 @@ Get node metadata: GET /v1/info ``` -Supported models and pricing are available at `/v1/models`. +Supported models and pricing are available at `/v1/models`. Upstream provider path discovery is available at `/v1/models/paths` and `/v1/models/paths/model?model_id=...`. ## Next Steps diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 40930676..45387fb4 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -137,6 +137,7 @@ Use environment variables for: | `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` | | `CORS_ORIGINS` | Allowed CORS origins | `*` | | `RELAYS` | Nostr relays (comma-separated) | (default set) | +| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to disable | `600` | ### Priority @@ -156,3 +157,8 @@ Manage which AI models you offer: - **Create aliases** — friendly names for models See [Pricing](pricing.md) for per-model pricing strategies. + +Model path discovery is refreshed in the background and exposed through +`/v1/models/paths`. The response groups each client-visible model ID with the +provider paths that may appear in chat-completion response metadata. Tune the +refresh cadence with `MODEL_PATHS_REFRESH_INTERVAL_SECONDS`. diff --git a/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py b/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py new file mode 100644 index 00000000..e5ea85f4 --- /dev/null +++ b/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py @@ -0,0 +1,66 @@ +"""add model_paths table + +Revision ID: d7e8f9a0b1c2 +Revises: c6d7e8f9a0b1 +Create Date: 2026-07-05 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "d7e8f9a0b1c2" +down_revision = "c6d7e8f9a0b1" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + if "model_paths" in inspector.get_table_names(): + return + + op.create_table( + "model_paths", + sa.Column("id", sa.Integer(), primary_key=True, autoincrement=True), + sa.Column("model_id", sa.String(), nullable=False), + sa.Column("path", sa.String(), nullable=False), + sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint( + ["upstream_provider_id"], + ["upstream_providers.id"], + ondelete="CASCADE", + ), + sa.UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + op.create_index( + "ix_model_paths_model_id", + "model_paths", + ["model_id"], + ) + op.create_index( + "ix_model_paths_upstream_provider_id", + "model_paths", + ["upstream_provider_id"], + ) + + +def downgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + if "model_paths" not in inspector.get_table_names(): + return + + existing_indexes = {idx["name"] for idx in inspector.get_indexes("model_paths")} + if "ix_model_paths_upstream_provider_id" in existing_indexes: + op.drop_index("ix_model_paths_upstream_provider_id", table_name="model_paths") + if "ix_model_paths_model_id" in existing_indexes: + op.drop_index("ix_model_paths_model_id", table_name="model_paths") + op.drop_table("model_paths") diff --git a/routstr/core/db.py b/routstr/core/db.py index 586f467e..4023c0c9 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -210,6 +210,41 @@ class ModelRow(SQLModel, table=True): # type: ignore upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") +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 + can therefore have several rows — one per direct provider path plus one per + OpenRouter sub-provider endpoint. + """ + + __tablename__ = "model_paths" + __table_args__ = ( + UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + id: int | None = Field(default=None, primary_key=True) + model_id: str = Field( + index=True, description="Client-visible /v1/models id (forwarded_model_id or id)" + ) + path: str = Field( + description="Provider path stamped on chat completion responses, e.g. " + "'anthropic' or 'openrouter:Anthropic'" + ) + upstream_provider_id: int = Field( + index=True, + foreign_key="upstream_providers.id", + ondelete="CASCADE", + description="upstream_providers.id this path was discovered from", + ) + + class LightningInvoice(SQLModel, table=True): # type: ignore __tablename__ = "lightning_invoices" diff --git a/routstr/core/main.py b/routstr/core/main.py index 903d6d1a..d105331a 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -58,6 +58,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: providers_task = None models_refresh_task = None model_maps_refresh_task = None + model_paths_refresh_task = None key_reset_task = None stale_reservation_task = None dead_key_prune_task = None @@ -127,6 +128,12 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: refresh_upstreams_models_periodically(get_upstreams) ) model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) + if global_settings.model_paths_refresh_interval_seconds > 0: + from ..upstream.model_paths import refresh_model_paths_periodically + + model_paths_refresh_task = asyncio.create_task( + refresh_model_paths_periodically(get_upstreams) + ) payout_task = asyncio.create_task(periodic_payout()) if global_settings.nsec: nip91_task = asyncio.create_task(announce_provider()) @@ -173,6 +180,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: models_refresh_task.cancel() if model_maps_refresh_task is not None: model_maps_refresh_task.cancel() + if model_paths_refresh_task is not None: + model_paths_refresh_task.cancel() if key_reset_task is not None: key_reset_task.cancel() if stale_reservation_task is not None: @@ -206,6 +215,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(models_refresh_task) if model_maps_refresh_task is not None: tasks_to_wait.append(model_maps_refresh_task) + if model_paths_refresh_task is not None: + tasks_to_wait.append(model_paths_refresh_task) if key_reset_task is not None: tasks_to_wait.append(key_reset_task) if stale_reservation_task is not None: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 3a144e10..5fce93af 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -94,6 +94,9 @@ class Settings(BaseSettings): models_refresh_interval_seconds: int = Field( default=360, env="MODELS_REFRESH_INTERVAL_SECONDS" ) + model_paths_refresh_interval_seconds: int = Field( + default=600, env="MODEL_PATHS_REFRESH_INTERVAL_SECONDS" + ) enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") diff --git a/routstr/payment/models.py b/routstr/payment/models.py index e6ea643d..b21669a8 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -596,6 +596,29 @@ async def test_model( } +@models_router.get("/v1/models/paths") +@models_router.get("/v1/models/paths/", include_in_schema=False) +async def model_paths() -> dict: + """All models with every upstream provider path they are reachable through.""" + from ..upstream.model_paths import get_all_model_paths + + return {"data": await get_all_model_paths()} + + +@models_router.get("/v1/models/paths/model") +@models_router.get("/v1/models/paths/model/", include_in_schema=False) +async def model_paths_for_model(model_id: str) -> dict: + """Paths for a single model. + + Uses a query parameter (``?model_id=...``) under a fully static route so + model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL + encoding and there is no dynamic-route ambiguity. + """ + from ..upstream.model_paths import get_paths_for_model + + return {"data": await get_paths_for_model(model_id)} + + @models_router.get("/v1/models") @models_router.get("/v1/models/", include_in_schema=False) @models_router.get("/models") diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py new file mode 100644 index 00000000..5ce1f527 --- /dev/null +++ b/routstr/upstream/model_paths.py @@ -0,0 +1,348 @@ +"""Model-path discovery service. + +Exposes every upstream provider path a Routstr model is reachable through. +This is discovery/visibility data only — routing still selects the cheapest or +best provider separately. + +A *path* is the provider string that may appear in Routstr chat completion +responses (see ``BaseUpstreamProvider._apply_provider_field``): + +- Direct upstream -> ```` e.g. ``anthropic`` +- Generic/custom OpenRouter-compatible upstream -> ``generic:`` +- Native OpenRouter routing to a sub-provider -> ``openrouter:`` + +Native OpenRouter does not emit a useful bare ``openrouter`` path when no +sub-provider is present; it reports ``unknown`` instead. +""" + +from __future__ import annotations + +import asyncio +import random +from typing import Callable + +import httpx +from sqlmodel import col, delete, select + +from ..core.db import ModelPathRow, create_session +from ..core.logging import get_logger +from .base import BaseUpstreamProvider + +logger = get_logger(__name__) + +# Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds +# of models does not open hundreds of concurrent requests every refresh. +_OPENROUTER_CONCURRENCY = 5 +_OPENROUTER_TIMEOUT_SECONDS = 10.0 + + +def is_openrouter_base_url(base_url: str | None) -> bool: + """True when ``base_url`` points at OpenRouter. + + Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``: + that predicate also returns True for native Anthropic (correct for + cache-control, wrong for OpenRouter endpoint discovery). This one keys only + on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched + while native Anthropic is not. + """ + return "openrouter.ai" in (base_url or "") + + +def exposed_model_id(model: object) -> str: + """Client-visible ``/v1/models`` id for a cached model.""" + forwarded = getattr(model, "forwarded_model_id", None) + return forwarded or getattr(model, "id") + + +def public_model_id(model_id: str) -> str: + """Model id exposed by model-path API responses. + + Provider-prefixed ids such as ``z-ai/glm-5v-turbo`` are returned as + ``glm-5v-turbo`` so clients can search and display the same unqualified id + they pass to ``/v1/models/paths/model``. + """ + return model_id.rsplit("/", 1)[-1] + + +def openrouter_author_slug(model: object) -> str | None: + """Return a canonical ``author/slug`` for the OpenRouter endpoints API. + + OpenRouter requires the canonical id, never ``forwarded_model_id``. Prefer + ``canonical_slug``, then a slash-containing ``id``; otherwise there is no + usable form and endpoint discovery is skipped for this model. + """ + canonical = getattr(model, "canonical_slug", None) + if canonical and "/" in canonical: + return canonical + model_id = getattr(model, "id", None) + if model_id and "/" in model_id: + return model_id + return None + + +async def _fetch_openrouter_endpoint_paths( + client: httpx.AsyncClient, + base_url: str, + api_key: str, + author_slug: str, + path_prefix: str, + semaphore: asyncio.Semaphore, +) -> list[str]: + """Return ``:`` paths for one model, or ``[]``. + + Failures (network, rate limit, bad payload) are logged and swallowed so one + model never breaks the whole refresh. + """ + url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" + headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} + async with semaphore: + try: + resp = await client.get( + url, headers=headers, timeout=_OPENROUTER_TIMEOUT_SECONDS + ) + except Exception as e: # noqa: BLE001 - isolate per-model failures + logger.warning( + "OpenRouter endpoint discovery request failed", + extra={"author_slug": author_slug, "error": str(e)}, + ) + return [] + + if resp.status_code == 429: + logger.warning( + "OpenRouter endpoint discovery rate-limited", + extra={"author_slug": author_slug}, + ) + return [] + if resp.status_code != 200: + logger.warning( + "OpenRouter endpoint discovery non-200", + extra={"author_slug": author_slug, "status_code": resp.status_code}, + ) + return [] + + try: + endpoints = resp.json().get("data", {}).get("endpoints", []) + except Exception as e: # noqa: BLE001 + logger.warning( + "OpenRouter endpoint discovery bad payload", + extra={"author_slug": author_slug, "error": str(e)}, + ) + return [] + + paths: list[str] = [] + for endpoint in endpoints: + provider_name = (endpoint or {}).get("provider_name") + if provider_name: + paths.append(f"{path_prefix}:{provider_name}") + # De-duplicate while preserving order. + return list(dict.fromkeys(paths)) + + +async def _collect_provider_paths( + upstream: BaseUpstreamProvider, +) -> list[tuple[str, str]]: + """Collect ``(model_id, path)`` pairs for one provider instance. + + Emits the direct ```` path for normal upstreams. For + OpenRouter-compatible providers, emits one path per OpenRouter sub-provider + endpoint, prefixed the same way response stamping prefixes it. + """ + provider_type = (upstream.provider_type or "").strip() + models = [m for m in upstream.get_cached_models() if getattr(m, "enabled", True)] + is_openrouter = is_openrouter_base_url(upstream.base_url) + + pairs: list[tuple[str, str]] = [] + if not is_openrouter: + for model in models: + if provider_type: + pairs.append((exposed_model_id(model), provider_type)) + return pairs + + if not provider_type: + return pairs + + semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) + async with httpx.AsyncClient() as client: + + async def _for_model(model: object) -> list[tuple[str, str]]: + author_slug = openrouter_author_slug(model) + if not author_slug: + return [] + paths = await _fetch_openrouter_endpoint_paths( + client, + upstream.base_url, + upstream.api_key, + author_slug, + provider_type, + semaphore, + ) + model_id = exposed_model_id(model) + return [(model_id, path) for path in paths] + + results = await asyncio.gather( + *(_for_model(m) for m in models), return_exceptions=True + ) + + for result in results: + if isinstance(result, BaseException): + logger.warning( + "OpenRouter endpoint discovery task errored", + extra={"provider": provider_type, "error": str(result)}, + ) + continue + pairs.extend(result) + + return pairs + + +async def _persist_provider_paths( + upstream_provider_id: int, pairs: list[tuple[str, str]] +) -> None: + """Replace all rows for ``upstream_provider_id`` with ``pairs``. + + Replacement (not upsert) so stale paths disappear when provider config or + upstream availability changes. + """ + unique_pairs = list(dict.fromkeys(pairs)) + async with create_session() as session: + await session.exec( # type: ignore[call-overload] + delete(ModelPathRow).where( + col(ModelPathRow.upstream_provider_id) == upstream_provider_id + ) + ) + for model_id, path in unique_pairs: + session.add( + ModelPathRow( + model_id=model_id, + path=path, + upstream_provider_id=upstream_provider_id, + ) + ) + await session.commit() + + +async def _prune_inactive_provider_paths(active_provider_ids: set[int]) -> None: + """Delete paths for providers no longer present in the live upstream set.""" + async with create_session() as session: + stmt = delete(ModelPathRow) + if active_provider_ids: + stmt = stmt.where( + col(ModelPathRow.upstream_provider_id).not_in(active_provider_ids) + ) + await session.exec(stmt) # type: ignore[call-overload] + await session.commit() + + +async def refresh_model_paths( + upstreams: list[BaseUpstreamProvider], +) -> None: + """Recompute and persist model paths for every enabled provider. + + One provider's failure is logged and isolated; it must not break the rest. + """ + active_provider_ids = { + upstream.db_id for upstream in upstreams if upstream.db_id is not None + } + await _prune_inactive_provider_paths(active_provider_ids) + + for upstream in upstreams: + if upstream.db_id is None: + continue + try: + pairs = await _collect_provider_paths(upstream) + await _persist_provider_paths(upstream.db_id, pairs) + except Exception as e: # noqa: BLE001 - isolate per-provider failures + logger.error( + "Failed to refresh model paths for provider", + extra={ + "provider": upstream.provider_type or upstream.base_url, + "db_id": upstream.db_id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + + +async def refresh_model_paths_periodically( + upstreams_provider: ( + Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider] + ), +) -> None: + """Background task mirroring ``refresh_upstreams_models_periodically``.""" + from ..core.settings import settings + + interval = getattr(settings, "model_paths_refresh_interval_seconds", 0) + if not interval or interval <= 0: + logger.info("Model paths refresh disabled (interval <= 0)") + return + + def _resolve_upstreams() -> list[BaseUpstreamProvider]: + if callable(upstreams_provider): + return upstreams_provider() + return upstreams_provider + + while True: + try: + await refresh_model_paths(_resolve_upstreams()) + except asyncio.CancelledError: + break + except Exception as e: # noqa: BLE001 + logger.error( + "Error in model paths refresh loop", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + try: + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) + except asyncio.CancelledError: + break + + +async def get_all_model_paths() -> list[dict]: + """All models with their paths, shaped for ``GET /v1/models/paths``.""" + async with create_session() as session: + rows = ( + await session.exec(select(ModelPathRow).order_by(ModelPathRow.model_id)) + ).all() + + grouped: dict[str, list[dict]] = {} + seen_paths: dict[str, set[str]] = {} + for row in rows: + model_id = public_model_id(row.model_id) + if row.path in seen_paths.setdefault(model_id, set()): + continue + seen_paths[model_id].add(row.path) + grouped.setdefault(model_id, []).append({"path": row.path}) + return [{"id": model_id, "paths": paths} for model_id, paths in grouped.items()] + + +async def get_paths_for_model(model_id: str) -> list[dict]: + """Paths for a single model, shaped for ``GET /v1/models/paths/model``. + + Match by the public, unqualified model id, mirroring the model cache alias + behavior. Both ``deepseek-v4-pro`` and ``deepseek/deepseek-v4-pro`` resolve + every row whose stored id has the same base model id. + """ + requested_id = public_model_id(model_id) + async with create_session() as session: + rows = ( + await session.exec( + select(ModelPathRow).order_by( + ModelPathRow.path, + col(ModelPathRow.upstream_provider_id), + ModelPathRow.model_id, + ) + ) + ).all() + + seen: set[str] = set() + paths: list[dict] = [] + for row in rows: + if public_model_id(row.model_id) != requested_id: + continue + if row.path in seen: + continue + seen.add(row.path) + paths.append({"path": row.path}) + return paths diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py new file mode 100644 index 00000000..75468383 --- /dev/null +++ b/tests/unit/test_model_paths.py @@ -0,0 +1,591 @@ +"""Tests for the model-path discovery service and endpoints.""" + +from __future__ import annotations + +import asyncio +import os +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Any, AsyncGenerator + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.db import UpstreamProviderRow # noqa: E402 +from routstr.payment.models import models_router # noqa: E402 +from routstr.upstream import model_paths as mp # noqa: E402 + +# --------------------------------------------------------------------------- # +# Fakes +# --------------------------------------------------------------------------- # + + +def _model( + id: str, + *, + forwarded_model_id: str | None = None, + canonical_slug: str | None = None, + enabled: bool = True, +) -> SimpleNamespace: + return SimpleNamespace( + id=id, + forwarded_model_id=forwarded_model_id, + canonical_slug=canonical_slug, + enabled=enabled, + ) + + +class _FakeProvider: + def __init__( + self, + *, + provider_type: str, + base_url: str, + models: list[SimpleNamespace], + db_id: int | None = 1, + api_key: str = "sk-test", + ) -> None: + self.provider_type = provider_type + self.base_url = base_url + self.api_key = api_key + self.db_id = db_id + self._models = models + + def get_cached_models(self) -> list[SimpleNamespace]: + return self._models + + +class _FakeResponse: + def __init__(self, status_code: int, payload: Any) -> None: + self.status_code = status_code + self._payload = payload + + def json(self) -> Any: + return self._payload + + +@pytest.fixture +async def patched_session( + monkeypatch: pytest.MonkeyPatch, +) -> AsyncGenerator[AsyncEngine, None]: + """Bind the service's ``create_session`` to a fresh in-memory engine.""" + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + # Seed the FK target so ModelPathRow inserts satisfy the constraint. + async with AsyncSession(engine) as session: + for pid in (1, 2): + session.add( + UpstreamProviderRow( + id=pid, + slug=f"p{pid}", + provider_type="anthropic" if pid == 1 else "openrouter", + base_url=f"https://provider-{pid}", + api_key=f"k{pid}", + ) + ) + await session.commit() + + @asynccontextmanager + async def _factory() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + monkeypatch.setattr(mp, "create_session", _factory) + yield engine + await engine.dispose() + + +# --------------------------------------------------------------------------- # +# Predicates / pure helpers +# --------------------------------------------------------------------------- # + + +def test_is_openrouter_base_url() -> None: + assert mp.is_openrouter_base_url("https://openrouter.ai/api/v1") is True + assert mp.is_openrouter_base_url("https://api.anthropic.com") is False + assert mp.is_openrouter_base_url(None) is False + + +def test_native_anthropic_not_openrouter() -> None: + """Native Anthropic must not be treated as OpenRouter-compatible even though + ``_upstream_accepts_cache_control`` returns True for it.""" + assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False + + +def test_exposed_model_id_prefers_forwarded() -> None: + assert ( + mp.exposed_model_id(_model("claude-x", forwarded_model_id="fwd-claude")) + == "fwd-claude" + ) + assert mp.exposed_model_id(_model("claude-x")) == "claude-x" + + +def test_public_model_id_strips_provider_prefix() -> None: + assert mp.public_model_id("z-ai/glm-5v-turbo") == "glm-5v-turbo" + assert mp.public_model_id("gpt-4o-mini") == "gpt-4o-mini" + + +def test_openrouter_author_slug_uses_canonical_not_forwarded() -> None: + m = _model( + "claude-opus-4.6", + forwarded_model_id="forwarded-only", + canonical_slug="anthropic/claude-opus-4.6", + ) + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_falls_back_to_slash_id() -> None: + m = _model("anthropic/claude-opus-4.6", canonical_slug="claude-opus-4.6") + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_none_when_no_slash() -> None: + m = _model("claude-opus-4.6", canonical_slug="claude-opus-4.6") + assert mp.openrouter_author_slug(m) is None + + +# --------------------------------------------------------------------------- # +# Collection +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_direct_provider_single_path_uses_provider_type() -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + ) + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("claude-opus-4.6", "anthropic")] + + +@pytest.mark.asyncio +async def test_direct_path_stores_exposed_model_id() -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("internal-id", forwarded_model_id="claude-opus-4.6")], + ) + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("claude-opus-4.6", "anthropic")] + + +@pytest.mark.asyncio +async def test_disabled_models_excluded() -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[ + _model("enabled-model"), + _model("disabled-model", enabled=False), + ], + ) + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("enabled-model", "anthropic")] + + +@pytest.mark.asyncio +async def test_openrouter_provider_adds_endpoint_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _fake_get( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + return _FakeResponse( + 200, + { + "data": { + "endpoints": [ + {"provider_name": "Anthropic"}, + {"provider_name": "Amazon Bedrock"}, + ] + } + }, + ) + + monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert ("claude-opus-4.6", "openrouter:Anthropic") in pairs + assert ("claude-opus-4.6", "openrouter:Amazon Bedrock") in pairs + assert ("claude-opus-4.6", "openrouter") not in pairs + assert len(pairs) == 2 + + +@pytest.mark.asyncio +async def test_generic_provider_with_openrouter_base_url_discovers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A generic provider pointed at OpenRouter exposes the response-stamped + ``generic:`` path, not a native ``openrouter:`` path.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _fake_get( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + return _FakeResponse( + 200, {"data": {"endpoints": [{"provider_name": "Anthropic"}]}} + ) + + monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("claude-opus-4.6", "generic:Anthropic")] + + +@pytest.mark.asyncio +async def test_openrouter_failure_degrades_gracefully( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _boom( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + raise RuntimeError("network down") + + monkeypatch.setattr("httpx.AsyncClient.get", _boom) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [] + + +@pytest.mark.asyncio +async def test_openrouter_rate_limit_skips_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _rate_limited( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + return _FakeResponse(429, {}) + + monkeypatch.setattr("httpx.AsyncClient.get", _rate_limited) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [] + + +@pytest.mark.asyncio +async def test_openrouter_fanout_is_bounded( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) + models = [ + _model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20) + ] + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=models, + ) + + state = {"current": 0, "max": 0} + + async def _slow_get( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + state["current"] += 1 + state["max"] = max(state["max"], state["current"]) + await asyncio.sleep(0.02) + state["current"] -= 1 + return _FakeResponse(200, {"data": {"endpoints": [{"provider_name": "X"}]}}) + + monkeypatch.setattr("httpx.AsyncClient.get", _slow_get) + + await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" + + +# --------------------------------------------------------------------------- # +# Persistence + query +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_refresh_replaces_stale_rows( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("m1", "anthropic"), ("m2", "anthropic")]) + first = await mp.get_all_model_paths() + assert {row["id"] for row in first} == {"m1", "m2"} + + # Second refresh with a different set — stale m2 must disappear. + await mp._persist_provider_paths(1, [("m1", "anthropic")]) + second = await mp.get_all_model_paths() + assert {row["id"] for row in second} == {"m1"} + + +@pytest.mark.asyncio +async def test_refresh_model_paths_prunes_inactive_provider_rows( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("m1", "anthropic")]) + await mp._persist_provider_paths(2, [("m2", "openrouter:Anthropic")]) + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + active_only = await mp.get_all_model_paths() + assert active_only == [{"id": "m1", "paths": [{"path": "anthropic"}]}] + + await mp.refresh_model_paths([]) + assert await mp.get_all_model_paths() == [] + + +@pytest.mark.asyncio +async def test_same_model_two_providers_two_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) + + data = await mp.get_all_model_paths() + assert len(data) == 1 + entry = data[0] + assert entry["id"] == "claude-opus-4.6" + paths = {p["path"] for p in entry["paths"]} + assert paths == {"anthropic", "openrouter:Anthropic"} + # No canonical_id anywhere. + assert "canonical_id" not in entry + assert all("canonical_id" not in p for p in entry["paths"]) + + +@pytest.mark.asyncio +async def test_get_all_model_paths_deduplicates_visible_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) + + assert await mp.get_all_model_paths() == [ + {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + ] + + +@pytest.mark.asyncio +async def test_get_all_model_paths_returns_unqualified_model_ids( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) + await mp._persist_provider_paths(5, [("openai/gpt-4o-mini", "openrouter:OpenAI")]) + + data = await mp.get_all_model_paths() + + assert {row["id"] for row in data} == {"glm-5v-turbo", "gpt-4o-mini"} + + +@pytest.mark.asyncio +async def test_get_paths_for_model_returns_only_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) + + paths = await mp.get_paths_for_model("claude-opus-4.6") + assert {p["path"] for p in paths} == {"anthropic", "openrouter:Anthropic"} + assert all(set(p.keys()) == {"path"} for p in paths) + assert await mp.get_paths_for_model("does-not-exist") == [] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) + + paths = await mp.get_paths_for_model("glm-5v-turbo") + + assert paths == [{"path": "openrouter:Z.AI"}] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_deduplicates_visible_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) + + assert await mp.get_paths_for_model("claude-opus-4.6") == [ + {"path": "anthropic"} + ] + assert await mp.get_paths_for_model("anthropic/claude-opus-4.6") == [ + {"path": "anthropic"} + ] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(7, [("deepseek-v4-pro", "generic")]) + await mp._persist_provider_paths( + 4, [("deepseek/deepseek-v4-pro", "openrouter:DeepSeek")] + ) + + short_paths = await mp.get_paths_for_model("deepseek-v4-pro") + prefixed_paths = await mp.get_paths_for_model("deepseek/deepseek-v4-pro") + + assert short_paths == [ + {"path": "generic"}, + {"path": "openrouter:DeepSeek"}, + ] + assert prefixed_paths == short_paths + + +@pytest.mark.asyncio +async def test_refresh_model_paths_skips_provider_without_db_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=None, + ) + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + assert await mp.get_all_model_paths() == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_isolates_provider_failure( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + good = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + bad = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("m", canonical_slug="a/m")], + db_id=2, + ) + + original = mp._collect_provider_paths + + async def _maybe_fail(upstream: Any) -> list[tuple[str, str]]: + if upstream is bad: + raise RuntimeError("boom") + return await original(upstream) # type: ignore[arg-type] + + monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) + + await mp.refresh_model_paths([good, bad]) # type: ignore[list-item] + data = await mp.get_all_model_paths() + assert {row["id"] for row in data} == {"claude-opus-4.6"} + + +# --------------------------------------------------------------------------- # +# Endpoints +# --------------------------------------------------------------------------- # + + +def _make_model_paths_app() -> FastAPI: + app = FastAPI() + app.include_router(models_router) + return app + + +def test_model_paths_endpoint_returns_all_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def _fake_get_all_model_paths() -> list[dict[str, Any]]: + return [ + { + "id": "claude-opus-4.6", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"}, + ], + } + ] + + monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) + + response = TestClient(_make_model_paths_app()).get("/v1/models/paths") + + assert response.status_code == 200 + assert response.json() == { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"}, + ], + } + ] + } + + +def test_model_paths_for_model_endpoint_accepts_slash_model_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[str] = [] + + async def _fake_get_paths_for_model(model_id: str) -> list[dict[str, Any]]: + calls.append(model_id) + return [{"path": "generic:Anthropic"}] + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", + params={"model_id": "anthropic/claude-opus-4.6"}, + ) + + assert response.status_code == 200 + assert response.json() == {"data": [{"path": "generic:Anthropic"}]} + assert calls == ["anthropic/claude-opus-4.6"] From dc25659cff29aad9316344c6c5cf848fe6e95181 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 7 Jul 2026 11:11:46 +0200 Subject: [PATCH 02/14] only activ model should be visible --- routstr/upstream/model_paths.py | 136 ++++++++++++++++++++++++++++++-- tests/unit/test_model_paths.py | 127 ++++++++++++++++++++++++++++- 2 files changed, 254 insertions(+), 9 deletions(-) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 5ce1f527..94d21833 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -19,15 +19,19 @@ from __future__ import annotations import asyncio import random -from typing import Callable +from typing import TYPE_CHECKING, Callable import httpx +from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select -from ..core.db import ModelPathRow, create_session +from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session from ..core.logging import get_logger from .base import BaseUpstreamProvider +if TYPE_CHECKING: + from ..payment.models import Model + logger = get_logger(__name__) # Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds @@ -138,8 +142,117 @@ async def _fetch_openrouter_endpoint_paths( return list(dict.fromkeys(paths)) +async def _load_model_visibility() -> tuple[ + dict[str, tuple[ModelRow, float]], set[str], set[int] +]: + """Load the same DB model visibility inputs used by routing. + + ``refresh_model_maps`` builds routing from enabled providers, enabled DB + override rows, and disabled model ids. Model-path discovery uses the same + view so the discovery API does not advertise models routing would hide and + reports forwarded aliases from DB overrides consistently with ``/v1/models``. + """ + async with create_session() as session: + query = select(UpstreamProviderRow).options( + selectinload(UpstreamProviderRow.models) # type: ignore[arg-type] + ) + provider_rows = (await session.exec(query)).all() + + overrides_by_id: dict[str, tuple[ModelRow, float]] = {} + disabled_model_ids: set[str] = set() + enabled_provider_ids: set[int] = set() + + for provider in provider_rows: + if not provider.enabled: + continue + if provider.id is not None: + enabled_provider_ids.add(provider.id) + for model in provider.models: + if model.enabled: + overrides_by_id[model.id] = (model, provider.provider_fee) + else: + disabled_model_ids.add(model.id) + + return overrides_by_id, disabled_model_ids, enabled_provider_ids + + +def _row_to_visible_model( + model_id: str, + row: ModelRow, + provider_fee: float, +) -> Model | None: + """Convert an enabled DB override row into a routed model object.""" + from ..payment.models import _row_to_model + + try: + return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) + except Exception as exc: # noqa: BLE001 - skip invalid override row + logger.warning( + "Skipping invalid model override while collecting model paths", + extra={ + "model_id": model_id, + "upstream_provider_id": getattr(row, "upstream_provider_id", None), + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) + return None + + +def _apply_model_visibility( + upstream: BaseUpstreamProvider, + overrides_by_id: dict[str, tuple[ModelRow, float]] | None, + disabled_model_ids: set[str] | None, +) -> list[object]: + """Return provider models after DB disabled/override state is applied.""" + overrides_by_id = overrides_by_id or {} + disabled_model_ids = disabled_model_ids or set() + visible_models: list[object] = [] + seen_model_ids: set[str] = set() + + for model in upstream.get_cached_models(): + model_id = getattr(model, "id", "") + if not getattr(model, "enabled", True) or model_id in disabled_model_ids: + continue + + if model_id in overrides_by_id: + override_row, provider_fee = overrides_by_id[model_id] + visible_model = _row_to_visible_model(model_id, override_row, provider_fee) + if visible_model is None: + continue + model = visible_model + + if not getattr(model, "enabled", True): + continue + visible_models.append(model) + seen_model_ids.add(model_id.lower()) + + upstream_provider_id = getattr(upstream, "db_id", None) + if isinstance(upstream_provider_id, int): + for model_id, (override_row, provider_fee) in overrides_by_id.items(): + if model_id in disabled_model_ids: + continue + if ( + getattr(override_row, "upstream_provider_id", None) + != upstream_provider_id + ): + continue + if model_id.lower() in seen_model_ids: + continue + override_model = _row_to_visible_model(model_id, override_row, provider_fee) + if override_model is None: + continue + if getattr(override_model, "enabled", True): + visible_models.append(override_model) + seen_model_ids.add(model_id.lower()) + + return visible_models + + async def _collect_provider_paths( upstream: BaseUpstreamProvider, + overrides_by_id: dict[str, tuple[ModelRow, float]] | None = None, + disabled_model_ids: set[str] | None = None, ) -> list[tuple[str, str]]: """Collect ``(model_id, path)`` pairs for one provider instance. @@ -148,7 +261,7 @@ async def _collect_provider_paths( endpoint, prefixed the same way response stamping prefixes it. """ provider_type = (upstream.provider_type or "").strip() - models = [m for m in upstream.get_cached_models() if getattr(m, "enabled", True)] + models = _apply_model_visibility(upstream, overrides_by_id, disabled_model_ids) is_openrouter = is_openrouter_base_url(upstream.base_url) pairs: list[tuple[str, str]] = [] @@ -240,16 +353,27 @@ async def refresh_model_paths( One provider's failure is logged and isolated; it must not break the rest. """ + ( + overrides_by_id, + disabled_model_ids, + enabled_provider_ids, + ) = await _load_model_visibility() active_provider_ids = { - upstream.db_id for upstream in upstreams if upstream.db_id is not None + upstream.db_id + for upstream in upstreams + if upstream.db_id is not None and upstream.db_id in enabled_provider_ids } await _prune_inactive_provider_paths(active_provider_ids) for upstream in upstreams: - if upstream.db_id is None: + if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: continue try: - pairs = await _collect_provider_paths(upstream) + pairs = await _collect_provider_paths( + upstream, + overrides_by_id=overrides_by_id, + disabled_model_ids=disabled_model_ids, + ) await _persist_provider_paths(upstream.db_id, pairs) except Exception as e: # noqa: BLE001 - isolate per-provider failures logger.error( diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 75468383..947c4312 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import json import os from contextlib import asynccontextmanager from types import SimpleNamespace @@ -18,7 +19,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") os.environ.setdefault("UPSTREAM_API_KEY", "test") -from routstr.core.db import UpstreamProviderRow # noqa: E402 +from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402 from routstr.payment.models import models_router # noqa: E402 from routstr.upstream import model_paths as mp # noqa: E402 @@ -42,6 +43,37 @@ def _model( ) +def _model_row( + id: str, + *, + upstream_provider_id: int = 1, + forwarded_model_id: str | None = None, + canonical_slug: str | None = None, + enabled: bool = True, +) -> ModelRow: + return ModelRow( + id=id, + upstream_provider_id=upstream_provider_id, + name=id, + created=0, + description="test model", + context_length=8192, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "test", + "instruct_type": None, + } + ), + pricing=json.dumps({"prompt": 0.000001, "completion": 0.000002}), + enabled=enabled, + forwarded_model_id=forwarded_model_id, + canonical_slug=canonical_slug, + ) + + class _FakeProvider: def __init__( self, @@ -360,6 +392,70 @@ async def test_refresh_replaces_stale_rows( assert {row["id"] for row in second} == {"m1"} +@pytest.mark.asyncio +async def test_refresh_model_paths_excludes_db_disabled_override( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("disabled-by-db", enabled=False)) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("disabled-by-db")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_uses_db_forwarded_alias( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("internal-id", forwarded_model_id="public-alias")) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("internal-id")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [ + {"id": "public-alias", "paths": [{"path": "anthropic"}]} + ] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cache( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("deployment-id", forwarded_model_id="public-deployment")) + await session.commit() + + provider = _FakeProvider( + provider_type="generic", + base_url="https://custom-provider/v1", + models=[], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [ + {"id": "public-deployment", "paths": [{"path": "generic"}]} + ] + + @pytest.mark.asyncio async def test_refresh_model_paths_prunes_inactive_provider_rows( patched_session: AsyncEngine, @@ -382,6 +478,29 @@ async def test_refresh_model_paths_prunes_inactive_provider_rows( assert await mp.get_all_model_paths() == [] +@pytest.mark.asyncio +async def test_refresh_model_paths_skips_disabled_db_provider( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("stale-model", "anthropic")]) + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.enabled = False + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("fresh-model")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [] + + @pytest.mark.asyncio async def test_same_model_two_providers_two_paths( patched_session: AsyncEngine, @@ -515,10 +634,12 @@ async def test_refresh_model_paths_isolates_provider_failure( original = mp._collect_provider_paths - async def _maybe_fail(upstream: Any) -> list[tuple[str, str]]: + async def _maybe_fail( + upstream: Any, *args: Any, **kwargs: Any + ) -> list[tuple[str, str]]: if upstream is bad: raise RuntimeError("boom") - return await original(upstream) # type: ignore[arg-type] + return await original(upstream, *args, **kwargs) # type: ignore[arg-type] monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) From ab80657507e921766dcba568be64b5d7d699fff2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 20:56:34 +0200 Subject: [PATCH 03/14] recreate model paths migration --- .../bda277ee4683_add_model_paths_table.py | 55 ++++++++++++++++ .../d7e8f9a0b1c2_add_model_paths_table.py | 66 ------------------- 2 files changed, 55 insertions(+), 66 deletions(-) create mode 100644 migrations/versions/bda277ee4683_add_model_paths_table.py delete mode 100644 migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py diff --git a/migrations/versions/bda277ee4683_add_model_paths_table.py b/migrations/versions/bda277ee4683_add_model_paths_table.py new file mode 100644 index 00000000..ab5ce131 --- /dev/null +++ b/migrations/versions/bda277ee4683_add_model_paths_table.py @@ -0,0 +1,55 @@ +"""add model paths table + +Revision ID: bda277ee4683 +Revises: fc4fa29630d2 +Create Date: 2026-07-24 20:54:56.822687 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "bda277ee4683" +down_revision = "fc4fa29630d2" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.create_table( + "model_paths", + sa.Column("id", sa.Integer(), nullable=False), + sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint( + ["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE" + ), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + op.create_index( + op.f("ix_model_paths_model_id"), "model_paths", ["model_id"], unique=False + ) + op.create_index( + op.f("ix_model_paths_upstream_provider_id"), + "model_paths", + ["upstream_provider_id"], + unique=False, + ) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths") + op.drop_index(op.f("ix_model_paths_model_id"), table_name="model_paths") + op.drop_table("model_paths") + # ### end Alembic commands ### diff --git a/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py b/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py deleted file mode 100644 index e5ea85f4..00000000 --- a/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py +++ /dev/null @@ -1,66 +0,0 @@ -"""add model_paths table - -Revision ID: d7e8f9a0b1c2 -Revises: c6d7e8f9a0b1 -Create Date: 2026-07-05 00:00:00.000000 -""" - -from __future__ import annotations - -import sqlalchemy as sa -from alembic import op - -revision = "d7e8f9a0b1c2" -down_revision = "c6d7e8f9a0b1" -branch_labels = None -depends_on = None - - -def upgrade() -> None: - conn = op.get_bind() - inspector = sa.inspect(conn) - if "model_paths" in inspector.get_table_names(): - return - - op.create_table( - "model_paths", - sa.Column("id", sa.Integer(), primary_key=True, autoincrement=True), - sa.Column("model_id", sa.String(), nullable=False), - sa.Column("path", sa.String(), nullable=False), - sa.Column("upstream_provider_id", sa.Integer(), nullable=False), - sa.ForeignKeyConstraint( - ["upstream_provider_id"], - ["upstream_providers.id"], - ondelete="CASCADE", - ), - sa.UniqueConstraint( - "model_id", - "path", - "upstream_provider_id", - name="uq_model_paths_model_path_provider", - ), - ) - op.create_index( - "ix_model_paths_model_id", - "model_paths", - ["model_id"], - ) - op.create_index( - "ix_model_paths_upstream_provider_id", - "model_paths", - ["upstream_provider_id"], - ) - - -def downgrade() -> None: - conn = op.get_bind() - inspector = sa.inspect(conn) - if "model_paths" not in inspector.get_table_names(): - return - - existing_indexes = {idx["name"] for idx in inspector.get_indexes("model_paths")} - if "ix_model_paths_upstream_provider_id" in existing_indexes: - op.drop_index("ix_model_paths_upstream_provider_id", table_name="model_paths") - if "ix_model_paths_model_id" in existing_indexes: - op.drop_index("ix_model_paths_model_id", table_name="model_paths") - op.drop_table("model_paths") From 4c292580e8d25d8239d828c8af7cc3448d1336d1 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 21:15:55 +0200 Subject: [PATCH 04/14] rebase model paths migration onto latest head --- ..._table.py => 4e0c3d195a49_add_model_paths_table.py} | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) rename migrations/versions/{bda277ee4683_add_model_paths_table.py => 4e0c3d195a49_add_model_paths_table.py} (91%) diff --git a/migrations/versions/bda277ee4683_add_model_paths_table.py b/migrations/versions/4e0c3d195a49_add_model_paths_table.py similarity index 91% rename from migrations/versions/bda277ee4683_add_model_paths_table.py rename to migrations/versions/4e0c3d195a49_add_model_paths_table.py index ab5ce131..61710f64 100644 --- a/migrations/versions/bda277ee4683_add_model_paths_table.py +++ b/migrations/versions/4e0c3d195a49_add_model_paths_table.py @@ -1,8 +1,8 @@ """add model paths table -Revision ID: bda277ee4683 -Revises: fc4fa29630d2 -Create Date: 2026-07-24 20:54:56.822687 +Revision ID: 4e0c3d195a49 +Revises: 7f2843d3f4e4 +Create Date: 2026-07-24 21:14:39.062179 """ import sqlalchemy as sa @@ -10,8 +10,8 @@ import sqlmodel from alembic import op # revision identifiers, used by Alembic. -revision = "bda277ee4683" -down_revision = "fc4fa29630d2" +revision = "4e0c3d195a49" +down_revision = "7f2843d3f4e4" branch_labels = None depends_on = None From 0a00527626202b5960c1152b66e8caa903b3f580 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:23:11 +0200 Subject: [PATCH 05/14] chore: apply ruff format repo-wide CI only runs ruff check, so format drift accumulated. Committed separately so the reformat noise stays out of functional commits. --- routstr/algorithm.py | 20 ++++++++-------- routstr/balance.py | 29 +++++++++++++++++++---- routstr/core/admin.py | 26 ++++++++++----------- routstr/core/log_manager.py | 21 ++++++++--------- routstr/core/usage_analytics_store.py | 29 +++++++++++------------ routstr/nostr/analytics.py | 8 +++++-- routstr/payment/cost_calculation.py | 19 ++++----------- routstr/payment/usage.py | 4 +--- routstr/upstream/azure.py | 4 +--- routstr/upstream/ehbp.py | 32 +++++++++++++++----------- routstr/upstream/gemini.py | 4 +--- routstr/upstream/gemini_messages.py | 4 +--- routstr/upstream/groq.py | 4 +++- routstr/upstream/litellm_routing.py | 4 +--- routstr/upstream/messages_dispatch.py | 12 +++------- routstr/upstream/ollama.py | 8 +++---- routstr/upstream/rate_limit.py | 4 +++- routstr/upstream/request_correction.py | 4 +--- routstr/upstream/routstr.py | 3 +-- routstr/upstream/xai.py | 4 +++- routstr/wallet.py | 13 ++++++++--- 21 files changed, 132 insertions(+), 124 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index fbc5388e..ef4b8574 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -232,7 +232,10 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: + if ( + model_to_use.forwarded_model_id + and model_to_use.forwarded_model_id not in aliases + ): aliases.append(model_to_use.forwarded_model_id) # Try to set each alias @@ -322,7 +325,10 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: + if ( + model_to_use.forwarded_model_id + and model_to_use.forwarded_model_id not in aliases + ): aliases.append(model_to_use.forwarded_model_id) for alias in aliases: @@ -342,16 +348,10 @@ def create_model_mappings( forwarded_model_ids, the one whose forwarded_model_id equals the requested alias wins. """ - if ( - model.forwarded_model_id - and model.forwarded_model_id.lower() == alias - ): + if model.forwarded_model_id and model.forwarded_model_id.lower() == alias: return 5 - if ( - model.id - and model.id.lower() == alias - ): + if model.id and model.id.lower() == alias: return 4 model_base = get_base_model_id(model.id) diff --git a/routstr/balance.py b/routstr/balance.py index 91b19ce5..03dc4d33 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -260,7 +260,11 @@ async def _lookup_key_no_create( async def _restore_balance( - session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str + session: AsyncSession, + hashed_key: str, + balance: int, + reserved_balance: int, + mint_url: str, ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -275,7 +279,11 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, + extra={ + "hashed_key": hashed_key, + "restored_balance": balance, + "mint_url": mint_url, + }, ) @@ -460,11 +468,23 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) error_msg = str(e) logger.error( "refund_wallet_endpoint: mint/send failed", @@ -685,7 +705,6 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} - @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 66a1d288..1521510e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -68,7 +68,9 @@ async def require_admin_api(request: Request) -> None: async with create_session() as session: result = await session.exec(select(CliToken).where(CliToken.token == token)) cli_token = result.first() - if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts): + if cli_token and ( + cli_token.expires_at is None or cli_token.expires_at > now_ts + ): cli_token.last_used_at = now_ts session.add(cli_token) await session.commit() @@ -255,16 +257,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) -> secret = await get_secret(session) if not secret.admin_password_hash: - raise HTTPException( - status_code=500, detail="Admin password not configured" - ) + raise HTTPException(status_code=500, detail="Admin password not configured") if not vault.verify_password( password_update.current_password, secret.admin_password_hash ): - raise HTTPException( - status_code=401, detail="Current password is incorrect" - ) + raise HTTPException(status_code=401, detail="Current password is incorrect") # Validate new password new_password = password_update.new_password.strip() @@ -980,9 +978,7 @@ async def update_upstream_provider_by_slug( lookup = _validate_slug(payload.slug) async with create_session() as session: result = await session.exec( - select(UpstreamProviderRow).where( - UpstreamProviderRow.slug == lookup - ) + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup) ) provider = result.first() if not provider: @@ -1669,7 +1665,11 @@ async def get_transactions_api( ) total = count_result.one() - stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit) + stmt = ( + base.order_by(col(CashuTransaction.created_at).desc()) + .offset(offset) + .limit(limit) + ) results = await session.exec(stmt) transactions = results.all() @@ -1679,9 +1679,7 @@ async def get_transactions_api( } -@admin_router.get( - "/api/lightning-invoices", dependencies=[Depends(require_admin_api)] -) +@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)]) async def get_lightning_invoices_api( status: str | None = None, purpose: str | None = None, diff --git a/routstr/core/log_manager.py b/routstr/core/log_manager.py index 0444dcbf..b111f68a 100644 --- a/routstr/core/log_manager.py +++ b/routstr/core/log_manager.py @@ -408,7 +408,9 @@ class LogManager: def get_error_details(self, hours: int = 24, limit: int = 100) -> dict: def compute() -> dict: try: - return self._usage_store.get_error_details(hours_back=hours, limit=limit) + return self._usage_store.get_error_details( + hours_back=hours, limit=limit + ) except Exception as e: logger.error( f"Usage analytics index failed, falling back to log scan: {e}" @@ -628,8 +630,7 @@ class LogManager: stats["total_tokens"] += input_tokens + output_tokens failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: stats["total_requests"] += 1 @@ -787,7 +788,9 @@ class LogManager: if bucket_key: model_mix_buckets[bucket_key][model] += 1 if revenue_msats > 0: - model_mix_revenue_buckets[bucket_key][model] += revenue_msats + model_mix_revenue_buckets[bucket_key][model] += ( + revenue_msats + ) model_mix_revenue_totals[model] += revenue_msats if input_tokens > 0 or output_tokens > 0: token_total = input_tokens + output_tokens @@ -801,8 +804,7 @@ class LogManager: bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: summary_stats["total_requests"] += 1 @@ -872,9 +874,7 @@ class LogManager: models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True) latest_errors = [ item - for _, item in sorted( - latest_errors_heap, key=lambda x: x[0], reverse=True - ) + for _, item in sorted(latest_errors_heap, key=lambda x: x[0], reverse=True) ] top_model_limit = max(1, min(model_limit, 20)) top_models_requests = [ @@ -1051,8 +1051,7 @@ class LogManager: bucket["warnings"] += 1 failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: bucket["total_requests"] += 1 diff --git a/routstr/core/usage_analytics_store.py b/routstr/core/usage_analytics_store.py index 7ba90e24..36fa4bcb 100644 --- a/routstr/core/usage_analytics_store.py +++ b/routstr/core/usage_analytics_store.py @@ -314,9 +314,7 @@ class UsageAnalyticsStore: if column in existing_columns: return - conn.execute( - f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}" - ) + conn.execute(f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}") logger.info(f"Migrated analytics schema: added {table}.{column}") def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None: @@ -364,7 +362,11 @@ class UsageAnalyticsStore: self._drop_index_tables_locked(conn) self._initialize_schema_locked(conn) - files = log_files if log_files is not None else sorted(self.logs_dir.glob("app_*.log")) + files = ( + log_files + if log_files is not None + else sorted(self.logs_dir.glob("app_*.log")) + ) for log_file in files: try: self._process_log_file_locked(conn, log_file, force_full_read=True) @@ -568,8 +570,7 @@ class UsageAnalyticsStore: model_bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: bucket["total_requests"] += 1 @@ -592,9 +593,9 @@ class UsageAnalyticsStore: if isinstance(max_cost, (int, float)) and max_cost > 0: max_cost_float = float(max_cost) bucket["refunds_msats"] += max_cost_float - model_updates[(minute_key, model)][ - "refunds_msats" - ] += max_cost_float + model_updates[(minute_key, model)]["refunds_msats"] += ( + max_cost_float + ) return ( end_offset, @@ -1032,7 +1033,9 @@ class UsageAnalyticsStore: """, (cutoff_timestamp,), ).fetchone() - total_error_count = int(total_error_count_row[0]) if total_error_count_row else 0 + total_error_count = ( + int(total_error_count_row[0]) if total_error_count_row else 0 + ) return { "errors": [ @@ -1204,11 +1207,7 @@ class UsageAnalyticsStore: total_successful = int(row["total_successful"]) total_revenue_msats = float(row["total_revenue_msats"]) total_tokens = int(row["total_tokens"]) - if ( - total_successful <= 0 - and total_revenue_msats <= 0 - and total_tokens <= 0 - ): + if total_successful <= 0 and total_revenue_msats <= 0 and total_tokens <= 0: continue bucket_ts = str(row["bucket_ts"]) diff --git a/routstr/nostr/analytics.py b/routstr/nostr/analytics.py index e568b5e0..8b6da590 100644 --- a/routstr/nostr/analytics.py +++ b/routstr/nostr/analytics.py @@ -215,7 +215,9 @@ def _build_window_payload( summary = dashboard.get("summary", {}) model_usage_mix = dashboard.get("model_usage_mix", {}) - summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {}) + summary_payload = _build_summary_payload( + summary if isinstance(summary, dict) else {} + ) usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {} top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload) @@ -338,7 +340,9 @@ async def publish_usage_analytics() -> None: nsec = (settings.nsec or "").strip() if not nsec: if not warned_missing_nsec: - logger.info("NSEC is not configured; skipping analytics sharing to Nostr") + logger.info( + "NSEC is not configured; skipping analytics sharing to Nostr" + ) warned_missing_nsec = True await asyncio.sleep(DISABLED_POLL_SECONDS) continue diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 37ac15d3..e7cee8ca 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -224,9 +224,7 @@ async def calculate_cost( "Token counts %s in the upstream response but cannot be " "priced; the request will appear in dashboards with the " "raw counts and a fixed max-cost charge.", - "are present" - if (input_tokens > 0 or output_tokens > 0) - else "are zero", + "are present" if (input_tokens > 0 or output_tokens > 0) else "are zero", extra={ "base_cost_msats": max_cost, "model": response_data.get("model", "unknown"), @@ -303,9 +301,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: # actually deducts from the balance. For non-BYOK providers (e.g. # OpenRouter) usage.cost already equals upstream_inference_cost, so we # fall through to the normal ``cost`` lookup below. - upstream_cost = _coerce_usd( - cost_details.get("upstream_inference_cost") - ) + upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost")) if upstream_cost > 0 and usage_data.get("is_byok"): byok_fee = _coerce_usd(usage_data.get("cost")) return upstream_cost + byok_fee @@ -336,8 +332,7 @@ def _get_pricing_rates( ``None`` means configured fixed pricing should be used by the caller. """ if settings.fixed_pricing and ( - settings.fixed_per_1k_input_tokens - or settings.fixed_per_1k_output_tokens + settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens ): return None @@ -393,12 +388,8 @@ def _get_pricing_rates( usd_per_sat = sats_usd_price() mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat - cache_read_usd = _coerce_usd( - pricing.get("cache_read_input_token_cost") - ) - cache_write_usd = _coerce_usd( - pricing.get("cache_creation_input_token_cost") - ) + cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost")) + cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost")) mscr_1k = ( cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat if cache_read_usd > 0 diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 02c90055..11d5c01e 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -110,9 +110,7 @@ def normalize_usage(usage_data: object) -> NormalizedUsage | None: if not isinstance(usage_data, dict): return None - output_tokens = _first_token_count( - usage_data, "completion_tokens", "output_tokens" - ) + output_tokens = _first_token_count(usage_data, "completion_tokens", "output_tokens") cache_read, cache_write = _extract_cache_tokens(usage_data) # ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index a693b763..985bcfd2 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -94,9 +94,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider): deployment_id = deployment_id.split("/")[-1] return f"openai/deployments/{deployment_id}/{clean_path}" - def get_request_base_url( - self, path: str, model_obj: "Model | None" = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: """Use endpoint root, stripping accidental /openai/v1 suffix if present.""" base_url = self.base_url.rstrip("/") marker = "/openai/v1" diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 96955492..a3d1505a 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -191,7 +191,9 @@ def _resolve_ehbp_target_url( otherwise the header is ignored so callers cannot redirect other providers or leak upstream API keys. """ - override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER + override_header = ( + profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER + ) if not override_header: return target_url enclave_url = _get_header_case_insensitive(headers, override_header) @@ -295,9 +297,7 @@ def _build_cost_info( return result -def _inject_cost_response_headers( - headers: dict[str, str], cost_info: dict -) -> None: +def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None: """Add per-request cost headers to an EHBP response. Since EHBP response bodies are opaque encrypted blobs, cost cannot be @@ -375,9 +375,7 @@ async def _compute_ehbp_actual_cost( resolved_upstream_model = ( actual_model_obj.forwarded_model_id or actual_model_obj.id ) - resolved_identity = _normalize_upstream_model_id( - resolved_upstream_model - ) + resolved_identity = _normalize_upstream_model_id(resolved_upstream_model) if resolved_identity != expected_identity: logger.info( "EHBP served model differs from requested, using actual " @@ -517,7 +515,9 @@ async def finalize_ehbp_actual_cost_payment( billing_key = await get_billing_key(key, session) key_hash = key.hashed_key billing_key_hash = billing_key.hashed_key - total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model))) + total_cost_msats = max( + 0, int(cost_info.get("total_msats", reserved_cost_for_model)) + ) now = int(time.time()) safe_reserved = case( @@ -560,7 +560,9 @@ async def finalize_ehbp_actual_cost_payment( ) child_result = await session.exec(child_stmt) # type: ignore[call-overload] - if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): + if result.rowcount == 0 or ( + child_result is not None and child_result.rowcount == 0 + ): await session.rollback() logger.error( "Failed to finalize EHBP usage-based payment", @@ -690,7 +692,9 @@ async def finalize_ehbp_max_cost_payment( else: child_result = None - if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): + if result.rowcount == 0 or ( + child_result is not None and child_result.rowcount == 0 + ): await session.rollback() logger.error( "Failed to finalize EHBP max-cost payment", @@ -1034,7 +1038,9 @@ async def forward_ehbp_x_cashu_request( target_url = _resolve_ehbp_target_url( target.url, path, headers, provider_type, profile ) - upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile) + upstream_headers = _prepare_ehbp_upstream_headers( + headers, target.headers, profile + ) request_body = await request.body() # Merge query params into the target URL @@ -1082,9 +1088,7 @@ async def forward_ehbp_x_cashu_request( usage_source = ( "header" if usage_header_name - and any( - k.lower() == usage_header_name.lower() for k, _ in resp.headers - ) + and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers) else ("trailer" if usage_header else "none") ) diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index 54de41a3..d58199c2 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -94,9 +94,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): """ return self.base_url.rstrip("/").removesuffix("/openai") + "/openai" - def get_request_base_url( - self, path: str, model_obj: "Model | None" = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: """Route every proxied request to the OpenAI-compat surface. Required because the stored ``base_url`` typically points at the diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index 11440b91..a7e41c9b 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -371,9 +371,7 @@ async def dispatch_gemini_messages( aggregates). """ if not request_body: - raise UpstreamError( - "Missing request body for /v1/messages", status_code=400 - ) + raise UpstreamError("Missing request body for /v1/messages", status_code=400) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index 17103c35..a0c9475e 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -20,7 +20,9 @@ class GroqUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": + def _build_from_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "GroqUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/upstream/litellm_routing.py b/routstr/upstream/litellm_routing.py index 0b2a92a3..b7790394 100644 --- a/routstr/upstream/litellm_routing.py +++ b/routstr/upstream/litellm_routing.py @@ -91,9 +91,7 @@ OLLAMA_HOST_HINTS: tuple[str, ...] = ( ) -def detect_litellm_prefix( - base_url: str | None, default: str = DEFAULT_PREFIX -) -> str: +def detect_litellm_prefix(base_url: str | None, default: str = DEFAULT_PREFIX) -> str: """Return the litellm provider prefix (`"/"`) for `base_url`. Falls back to `default` when the host doesn't match any known provider. diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 3d689922..efcf591b 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -108,9 +108,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: return events, buffer -def events_from_chunk( - chunk: object, sse_buffer: bytes -) -> tuple[list[dict], bytes]: +def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]: """Normalize a stream chunk into one or more event dicts. ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE @@ -201,9 +199,7 @@ async def aggregate_anthropic_events_to_message( raw_json = partial_json.pop(idx, None) if raw_json is not None and idx < len(blocks): try: - blocks[idx]["input"] = ( - json.loads(raw_json) if raw_json else {} - ) + blocks[idx]["input"] = json.loads(raw_json) if raw_json else {} except json.JSONDecodeError: blocks[idx]["input"] = raw_json elif etype == "message_delta": @@ -445,9 +441,7 @@ async def dispatch_anthropic_messages( on bad input or upstream failure. """ if not request_body: - raise UpstreamError( - "Missing request body for /v1/messages", status_code=400 - ) + raise UpstreamError("Missing request body for /v1/messages", status_code=400) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index 9fed0154..c4873ea0 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - def get_request_base_url( - self, path: str, model_obj: Model | None = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" return f"{self.base_url.rstrip('/')}/v1" @@ -185,7 +183,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): except Exception: self._models_cache = models_with_fees - self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache} + self._models_by_id = { + m.forwarded_model_id or m.id: m for m in self._models_cache + } logger.info( f"Refreshed models cache for {self.base_url}", extra={"model_count": len(models)}, diff --git a/routstr/upstream/rate_limit.py b/routstr/upstream/rate_limit.py index dca1ba5b..ac1eff78 100644 --- a/routstr/upstream/rate_limit.py +++ b/routstr/upstream/rate_limit.py @@ -119,7 +119,9 @@ def classify_rate_limit( retry_match = _RETRY_RE.search(redacted) if retry_match is not None: value = float(retry_match.group(1)) - retry_after = value / 1000.0 if retry_match.group(2).lower() == "ms" else value + retry_after = ( + value / 1000.0 if retry_match.group(2).lower() == "ms" else value + ) limit_name_match = _LIMIT_NAME_RE.search(redacted) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index c2ea5b1d..8e4379a0 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -84,9 +84,7 @@ def extract_error_message(response: Response) -> str: return "" -def strip_unsupported_param( - body: dict, error_message: str -) -> tuple[dict, str] | None: +def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: """Drop a top-level param the upstream named as unsupported/deprecated. Returns ``(new_body, param)`` (a new dict, original untouched) when the diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index 0371946a..de1aa3bd 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -50,8 +50,7 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): def normalize_request_path( self, path: str, model_obj: "Model | None" = None ) -> str: - """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr. - """ + """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.""" return path.lstrip("/") @classmethod diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 58caaba0..12e3dd93 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -21,7 +21,9 @@ class XAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": + def _build_from_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "XAIUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/wallet.py b/routstr/wallet.py index dd92d913..cef6902d 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -243,7 +243,9 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { - f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id) + f"{k.mint_url}/{k.unit.name}": sum( + p.amount for p in wallet.proofs if p.id == k.id + ) for k in wallet.keysets.values() } # Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet @@ -598,11 +600,16 @@ async def swap_to_primary_mint( # advance the counter so the next request derives fresh secrets. logger.warning( "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", - extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, + extra={ + "mint_quote_id": mint_quote.quote, + "minted_amount": minted_amount, + }, ) try: for keyset_id in primary_wallet.keysets: - await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) + await primary_wallet.restore_tokens_for_keyset( + keyset_id, to=1, batch=25 + ) await primary_wallet.load_proofs(reload=True) post_recovery_balance = primary_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance From f96acbb99cc5dbc11e64f9f933d8520031aaf059 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:23:31 +0200 Subject: [PATCH 06/14] fix: address model-paths review findings Provider scoping (items 1/2/6): - Key visibility maps on (model_id.lower(), upstream_provider_id), matching refresh_model_maps, so a disable/override row on one provider never leaks onto another provider's model, and matching is case-insensitive. Data safety (items 3/5): - Degraded OpenRouter fetches (network error, 429, non-200, bad payload) return None (unknown) instead of []; a provider whose path set is unknown keeps its previously persisted rows instead of being wiped. - Endpoint payload parsing moved fully inside try, with a list guard, so endpoints:null or non-list shapes are swallowed as documented. - refresh with an empty live upstream list is a no-op; the unfiltered DELETE in the prune path is gone (prune now keys off enabled DB rows). Hot path (items 4/12/14): - Persist uses chunked bulk INSERTs (one statement per 500 rows) instead of per-row ORM adds; redundant ix_model_paths_model_id index dropped. - Read routes filter in SQL instead of materializing the whole table, and output ordering is deterministic (public id + path), independent of rowid. - Visibility no longer rebuilds fully priced Model objects per override row; it reads id/forwarded_model_id/canonical_slug straight off ModelRow. Path/id contract (items 7/8/9/11): - discovery_path_for_subprovider/discovery_base_paths hooks on BaseUpstreamProvider, overridden by OpenRouterUpstreamProvider, mirror _apply_provider_field so discovery and response stamping cannot drift (openrouter:OpenRouter now correctly maps to unknown). - openrouter_author_slug falls back to a slash-containing forwarded_model_id, so admin-created alias rows are discoverable. - public_model_id splits on the first slash, same as get_base_model_id, so discovery ids can be sent to chat completions verbatim. Lifecycle (items 10/13): - ENABLE_MODEL_PATHS_REFRESH kill switch; interval and flag re-read every loop iteration, and the task idles (not exits) while disabled. - First 429 latches and aborts the remaining fan-out for the cycle; a per-cycle cache dedupes fetches across providers sharing a base URL. - refresh_model_maps prunes paths of disabled/deleted providers so admin mutations take effect immediately; rows carry updated_at and both endpoints expose it. Tests (item 15) rewritten through the public refresh entry point with transport-level httpx.MockTransport fakes, FK enforcement on, and coverage for the periodic loop. Migration re-chained onto 9c4d8e2f1a6b. --- docs/api/endpoints.md | 11 +- docs/provider/configuration.md | 3 +- .../4e0c3d195a49_add_model_paths_table.py | 15 +- routstr/core/db.py | 25 +- routstr/core/main.py | 23 +- routstr/core/settings.py | 18 +- routstr/payment/models.py | 12 +- routstr/proxy.py | 13 + routstr/upstream/base.py | 22 + routstr/upstream/model_paths.py | 469 ++++--- routstr/upstream/openrouter.py | 17 + tests/unit/test_fee_payout_migration.py | 4 +- tests/unit/test_model_paths.py | 1103 ++++++++++++----- 13 files changed, 1183 insertions(+), 552 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index efe57a68..44e6207e 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -374,10 +374,13 @@ GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 } ``` -Model IDs in responses are unqualified display IDs: provider prefixes such as -`z-ai/` or `openai/` are stripped. Path values match the provider string stamped -on chat-completion responses, such as `anthropic`, `generic:my-upstream`, or -`openrouter:Anthropic`. +Model IDs in responses are base model IDs: the leading provider prefix such as +`z-ai/` or `openai/` is stripped (the same rule routing uses, so the ID can be +sent back to `/v1/chat/completions` verbatim). Path values match the provider +string stamped on chat-completion responses, such as `anthropic`, +`generic:Anthropic`, `openrouter:Anthropic`, or `unknown` (native OpenRouter +with no usable sub-provider). Responses also carry an `updated_at` Unix +timestamp of the last successful refresh (`null` when no refresh has run). ## Wallet Management diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 6dd9d338..8bcbca9b 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -142,7 +142,8 @@ Use environment variables for: | `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` | | `CORS_ORIGINS` | Allowed CORS origins | `*` | | `RELAYS` | Nostr relays (comma-separated) | (default set) | -| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to disable | `600` | +| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to pause the refresh (previously discovered paths keep being served) | `600` | +| `ENABLE_MODEL_PATHS_REFRESH` | Kill switch for the background model-path refresh (OpenRouter endpoint fan-out) | `true` | ### Priority diff --git a/migrations/versions/4e0c3d195a49_add_model_paths_table.py b/migrations/versions/4e0c3d195a49_add_model_paths_table.py index 61710f64..5641e488 100644 --- a/migrations/versions/4e0c3d195a49_add_model_paths_table.py +++ b/migrations/versions/4e0c3d195a49_add_model_paths_table.py @@ -1,7 +1,7 @@ """add model paths table Revision ID: 4e0c3d195a49 -Revises: 7f2843d3f4e4 +Revises: 9c4d8e2f1a6b Create Date: 2026-07-24 21:14:39.062179 """ @@ -11,19 +11,19 @@ from alembic import op # revision identifiers, used by Alembic. revision = "4e0c3d195a49" -down_revision = "7f2843d3f4e4" +down_revision = "9c4d8e2f1a6b" branch_labels = None depends_on = None def upgrade() -> None: - # ### commands auto generated by Alembic - please adjust! ### op.create_table( "model_paths", sa.Column("id", sa.Integer(), nullable=False), sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False), sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.Column("updated_at", sa.Integer(), nullable=False, server_default="0"), sa.ForeignKeyConstraint( ["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE" ), @@ -35,21 +35,16 @@ def upgrade() -> None: name="uq_model_paths_model_path_provider", ), ) - op.create_index( - op.f("ix_model_paths_model_id"), "model_paths", ["model_id"], unique=False - ) + # No standalone index on model_id: the unique constraint's autoindex already + # leads on model_id. op.create_index( op.f("ix_model_paths_upstream_provider_id"), "model_paths", ["upstream_provider_id"], unique=False, ) - # ### end Alembic commands ### def downgrade() -> None: - # ### commands auto generated by Alembic - please adjust! ### op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths") - op.drop_index(op.f("ix_model_paths_model_id"), table_name="model_paths") op.drop_table("model_paths") - # ### end Alembic commands ### diff --git a/routstr/core/db.py b/routstr/core/db.py index 3cd34dcc..8b3586ec 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -222,7 +222,10 @@ async def release_stale_reservations( if released: logger.warning( "Released stale reservations", - extra={"released_reservations": released, "max_age_seconds": max_age_seconds}, + extra={ + "released_reservations": released, + "max_age_seconds": max_age_seconds, + }, ) return released @@ -255,9 +258,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in .where(col(ApiKey.total_spent) == 0) .where(col(ApiKey.total_requests) == 0) .where(col(ApiKey.parent_key_hash).is_(None)) - .where( - (col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff) - ) + .where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)) .where(~pending_invoice) .where(~has_children) ) @@ -331,8 +332,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore ), ) id: int | None = Field(default=None, primary_key=True) + # No standalone index on model_id: the unique constraint's autoindex already + # leads on model_id, so a second index only adds write amplification. model_id: str = Field( - index=True, description="Client-visible /v1/models id (forwarded_model_id or id)" + description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( description="Provider path stamped on chat completion responses, e.g. " @@ -344,6 +347,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore ondelete="CASCADE", description="upstream_providers.id this path was discovered from", ) + updated_at: int = Field( + default=0, + description="Unix timestamp of the refresh cycle that wrote this row", + ) class LightningInvoice(SQLModel, table=True): # type: ignore @@ -631,9 +638,7 @@ class CliToken(SQLModel, table=True): # type: ignore """Long-lived authorization token for CLI/agent use against admin endpoints.""" __tablename__ = "cli_tokens" - id: str = Field( - primary_key=True, default_factory=lambda: uuid.uuid4().hex - ) + id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex) token: str = Field(unique=True, index=True, description="Bearer token value") name: str = Field(description="Human-readable label for this token") created_at: int = Field(default_factory=lambda: int(time.time())) @@ -732,9 +737,7 @@ async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool: return result.rowcount == 1 -async def complete_routstr_fee_payout( - session: AsyncSession, paid_msats: int -) -> bool: +async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) -> bool: """Mark a checkpointed payout complete after the external payment succeeds.""" stmt = ( update(RoutstrFee) diff --git a/routstr/core/main.py b/routstr/core/main.py index 6eca6b52..979f5cdb 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -131,12 +131,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: refresh_upstreams_models_periodically(get_upstreams) ) model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) - if global_settings.model_paths_refresh_interval_seconds > 0: - from ..upstream.model_paths import refresh_model_paths_periodically + # Always started: the loop re-reads the enable flag and interval every + # iteration, so 0 -> N (or re-enabling) takes effect without a restart. + from ..upstream.model_paths import refresh_model_paths_periodically - model_paths_refresh_task = asyncio.create_task( - refresh_model_paths_periodically(get_upstreams) - ) + model_paths_refresh_task = asyncio.create_task( + refresh_model_paths_periodically(get_upstreams) + ) payout_task = asyncio.create_task(periodic_payout()) if global_settings.nsec: nip91_task = asyncio.create_task(announce_provider()) @@ -144,9 +145,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: if global_settings.providers_refresh_interval_seconds > 0: providers_task = asyncio.create_task(providers_cache_refresher()) key_reset_task = asyncio.create_task(periodic_key_reset()) - stale_reservation_task = asyncio.create_task( - periodic_stale_reservation_sweep() - ) + stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep()) dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) @@ -256,9 +255,7 @@ class _ImmutableStaticFiles(StaticFiles): async def get_response(self, path: str, scope: Scope) -> StarletteResponse: response = await super().get_response(path, scope) if response.status_code == 200: - response.headers["Cache-Control"] = ( - "public, max-age=31536000, immutable" - ) + response.headers["Cache-Control"] = "public, max-age=31536000, immutable" return response @@ -332,9 +329,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): # Serve the App Router RSC payload for the home page. @app.get("/index.txt", include_in_schema=False) async def serve_root_rsc() -> FileResponse: - return FileResponse( - UI_DIST_PATH / "index.txt", media_type="text/x-component" - ) + return FileResponse(UI_DIST_PATH / "index.txt", media_type="text/x-component") # Next.js is built with `trailingSlash: true`, so all UI page URLs end # with a slash (e.g. `/login/`). The proxy router catches `/{path:path}` diff --git a/routstr/core/settings.py b/routstr/core/settings.py index f5d0d2fd..ab91c3a4 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -100,8 +100,13 @@ class Settings(BaseSettings): ) enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") + enable_model_paths_refresh: bool = Field( + default=True, env="ENABLE_MODEL_PATHS_REFRESH" + ) refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") - refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS") + refund_sweep_ttl_seconds: int = Field( + default=604800, env="REFUND_SWEEP_TTL_SECONDS" + ) # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") @@ -120,9 +125,8 @@ class Settings(BaseSettings): # Discovery relays: list[str] = Field(default_factory=list, env="RELAYS") - enable_analytics_sharing: bool = Field( - default=True, env="ENABLE_ANALYTICS_SHARING" - ) + enable_analytics_sharing: bool = Field(default=True, env="ENABLE_ANALYTICS_SHARING") + def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]: """Discard unknown keys from persisted settings.""" @@ -330,7 +334,11 @@ class SettingsService: valid_fields = set(env_resolved.dict().keys()) merged_dict: dict[str, Any] = dict(env_resolved.dict()) merged_dict.update( - {k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields} + { + k: v + for k, v in db_json.items() + if v not in (None, "", [], {}) and k in valid_fields + } ) merged_dict = Settings(**merged_dict).dict() diff --git a/routstr/payment/models.py b/routstr/payment/models.py index b1e018d2..c433ddfa 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -455,7 +455,9 @@ async def _update_sats_pricing_once() -> None: for m in upstream.get_cached_models() ] upstream._models_cache = updated_models - upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models} + upstream._models_by_id = { + m.forwarded_model_id or m.id: m for m in updated_models + } updated_count += len(updated_models) if updated_count > 0: @@ -510,9 +512,7 @@ class ModelTestRequest(V2BaseModel): request_data: dict -@models_router.post( - "/api/models/test", dependencies=[Depends(_require_admin_api)] -) +@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) async def test_model( payload: ModelTestRequest, session: AsyncSession = Depends(get_session), @@ -601,7 +601,7 @@ async def model_paths() -> dict: """All models with every upstream provider path they are reachable through.""" from ..upstream.model_paths import get_all_model_paths - return {"data": await get_all_model_paths()} + return await get_all_model_paths() @models_router.get("/v1/models/paths/model") @@ -615,7 +615,7 @@ async def model_paths_for_model(model_id: str) -> dict: """ from ..upstream.model_paths import get_paths_for_model - return {"data": await get_paths_for_model(model_id)} + return await get_paths_for_model(model_id) @models_router.get("/v1/models") diff --git a/routstr/proxy.py b/routstr/proxy.py index c0b8794f..d3e9b59f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -183,6 +183,19 @@ async def refresh_model_maps() -> None: disabled_model_keys=disabled_model_keys, ) + # Keep model-path discovery in sync with admin mutations: disabling or + # deleting a provider must stop advertising its paths immediately rather + # than after the next timed refresh. + from .upstream.model_paths import prune_model_paths_for_inactive_providers + + try: + await prune_model_paths_for_inactive_providers() + except Exception as e: # noqa: BLE001 - discovery sync must not break routing + logger.warning( + "Failed to prune model paths for inactive providers", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + async def refresh_model_maps_periodically() -> None: """Background task to refresh model maps every minute.""" diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index a8dba7e3..5f9116c9 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -237,6 +237,28 @@ class BaseUpstreamProvider: except (TypeError, ValueError): pass + def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: + """Discovery path for a reported sub-provider name. + + Must produce exactly the value ``_apply_provider_field`` would stamp on + a response whose upstream payload reported ``sub_provider``, so the + model-path discovery API never advertises a path that cannot appear on + the wire. Subclasses that override ``_apply_provider_field`` must + override this to match. + """ + provider_type = (self.provider_type or "").strip() + if not provider_type: + return None + sub = (sub_provider or "").strip() + if not sub or sub == provider_type or sub.startswith(f"{provider_type}:"): + return sub or provider_type + return f"{provider_type}:{sub}" + + def discovery_base_paths(self) -> list[str]: + """Paths stamped when the upstream reports no sub-provider of its own.""" + provider_type = (self.provider_type or "").strip() + return [provider_type] if provider_type else [] + def _apply_provider_field(self, response_json: object) -> None: """Stamp the routstr ``provider`` field onto an upstream response payload. diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 94d21833..beaafa6a 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -5,32 +5,34 @@ This is discovery/visibility data only — routing still selects the cheapest or best provider separately. A *path* is the provider string that may appear in Routstr chat completion -responses (see ``BaseUpstreamProvider._apply_provider_field``): +responses. The strings emitted here are produced by the provider's own +``discovery_path_for_subprovider`` / ``discovery_base_paths`` hooks, which +mirror ``_apply_provider_field`` so discovery and response stamping cannot +drift: - Direct upstream -> ```` e.g. ``anthropic`` - Generic/custom OpenRouter-compatible upstream -> ``generic:`` - Native OpenRouter routing to a sub-provider -> ``openrouter:`` - -Native OpenRouter does not emit a useful bare ``openrouter`` path when no -sub-provider is present; it reports ``unknown`` instead. +- Native OpenRouter with no usable sub-provider -> ``unknown`` """ from __future__ import annotations import asyncio import random +import time from typing import TYPE_CHECKING, Callable import httpx +from sqlalchemy import insert, or_ from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session from ..core.logging import get_logger -from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..payment.models import Model + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -39,6 +41,20 @@ logger = get_logger(__name__) _OPENROUTER_CONCURRENCY = 5 _OPENROUTER_TIMEOUT_SECONDS = 10.0 +# Rows inserted per statement during persist. Keeps each INSERT bounded while +# avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle. +_PERSIST_CHUNK_SIZE = 500 + +# Visibility key used across this module: routing carries the provider +# dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)), +# so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps. +ModelKey = tuple[str, int] + + +def _make_http_client() -> httpx.AsyncClient: + """Client factory, separated so tests can substitute a mock transport.""" + return httpx.AsyncClient() + def is_openrouter_base_url(base_url: str | None) -> bool: """True when ``base_url`` points at OpenRouter. @@ -61,19 +77,21 @@ def exposed_model_id(model: object) -> str: def public_model_id(model_id: str) -> str: """Model id exposed by model-path API responses. - Provider-prefixed ids such as ``z-ai/glm-5v-turbo`` are returned as - ``glm-5v-turbo`` so clients can search and display the same unqualified id - they pass to ``/v1/models/paths/model``. + Uses the same rule as ``create_model_mappings.get_base_model_id`` and + ``resolve_model_alias`` — strip everything before the *first* slash — so + the id shown here can be sent back to ``/v1/chat/completions`` verbatim. """ - return model_id.rsplit("/", 1)[-1] + return model_id.split("/", 1)[1] if "/" in model_id else model_id def openrouter_author_slug(model: object) -> str | None: """Return a canonical ``author/slug`` for the OpenRouter endpoints API. - OpenRouter requires the canonical id, never ``forwarded_model_id``. Prefer - ``canonical_slug``, then a slash-containing ``id``; otherwise there is no - usable form and endpoint discovery is skipped for this model. + Prefer ``canonical_slug``, then a slash-containing ``id``, then a + slash-containing ``forwarded_model_id``. The forwarded id is exactly what + the proxy sends upstream for admin-created alias rows (``base.py`` forwards + ``forwarded_model_id or id``), so it is a valid OpenRouter id when the + bare ``id`` is a local alias with no slash. """ canonical = getattr(model, "canonical_slug", None) if canonical and "/" in canonical: @@ -81,24 +99,51 @@ def openrouter_author_slug(model: object) -> str | None: model_id = getattr(model, "id", None) if model_id and "/" in model_id: return model_id + forwarded = getattr(model, "forwarded_model_id", None) + if forwarded and "/" in forwarded: + return forwarded return None -async def _fetch_openrouter_endpoint_paths( +class _RefreshCycleState: + """Per-refresh shared state: fetch dedupe cache and rate-limit latch. + + ``endpoint_cache`` dedupes byte-identical ``/endpoints`` fetches when two + providers point at the same OpenRouter base URL. ``rate_limited`` latches + on the first 429 so the rest of the cycle stops hammering a throttled API; + the whole provider result then degrades to "unknown" instead of an empty + list, which preserves previously persisted rows. + """ + + def __init__(self) -> None: + self.endpoint_cache: dict[tuple[str, str], list[str] | None] = {} + self.rate_limited = False + + +async def _fetch_openrouter_endpoint_subproviders( client: httpx.AsyncClient, base_url: str, api_key: str, author_slug: str, - path_prefix: str, semaphore: asyncio.Semaphore, -) -> list[str]: - """Return ``:`` paths for one model, or ``[]``. + cycle: _RefreshCycleState, +) -> list[str] | None: + """Return sub-provider names for one model, or ``None`` when unknown. - Failures (network, rate limit, bad payload) are logged and swallowed so one - model never breaks the whole refresh. + ``None`` (not ``[]``) signals a degraded fetch — network failure, rate + limit, non-200, or an unparseable payload — so callers can distinguish + "this model has no endpoints" from "we could not find out". Failures are + logged and swallowed so one model never breaks the whole refresh. """ + cache_key = (base_url, author_slug) + if cache_key in cycle.endpoint_cache: + return cycle.endpoint_cache[cache_key] + if cycle.rate_limited: + return None + url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} + result: list[str] | None async with semaphore: try: resp = await client.get( @@ -109,48 +154,60 @@ async def _fetch_openrouter_endpoint_paths( "OpenRouter endpoint discovery request failed", extra={"author_slug": author_slug, "error": str(e)}, ) - return [] + cycle.endpoint_cache[cache_key] = None + return None if resp.status_code == 429: logger.warning( - "OpenRouter endpoint discovery rate-limited", + "OpenRouter endpoint discovery rate-limited; aborting cycle", extra={"author_slug": author_slug}, ) - return [] + cycle.rate_limited = True + cycle.endpoint_cache[cache_key] = None + return None if resp.status_code != 200: logger.warning( "OpenRouter endpoint discovery non-200", extra={"author_slug": author_slug, "status_code": resp.status_code}, ) - return [] + cycle.endpoint_cache[cache_key] = None + return None try: endpoints = resp.json().get("data", {}).get("endpoints", []) + if not isinstance(endpoints, list): + endpoints = [] + names: list[str] = [] + for endpoint in endpoints: + provider_name = ( + endpoint.get("provider_name") if isinstance(endpoint, dict) else None + ) + if provider_name: + names.append(provider_name) + result = list(dict.fromkeys(names)) except Exception as e: # noqa: BLE001 logger.warning( "OpenRouter endpoint discovery bad payload", extra={"author_slug": author_slug, "error": str(e)}, ) - return [] + result = None - paths: list[str] = [] - for endpoint in endpoints: - provider_name = (endpoint or {}).get("provider_name") - if provider_name: - paths.append(f"{path_prefix}:{provider_name}") - # De-duplicate while preserving order. - return list(dict.fromkeys(paths)) + cycle.endpoint_cache[cache_key] = result + return result async def _load_model_visibility() -> tuple[ - dict[str, tuple[ModelRow, float]], set[str], set[int] + dict[ModelKey, ModelRow], set[ModelKey], set[int] ]: """Load the same DB model visibility inputs used by routing. ``refresh_model_maps`` builds routing from enabled providers, enabled DB - override rows, and disabled model ids. Model-path discovery uses the same - view so the discovery API does not advertise models routing would hide and - reports forwarded aliases from DB overrides consistently with ``/v1/models``. + override rows, and disabled model keys — all keyed on + ``(model_id.lower(), upstream_provider_id)`` because ``ModelRow``'s primary + key is composite and the same id legitimately exists on several providers. + Model-path discovery uses the same keying so disabling a model on one + provider never hides it on another, and one provider's + ``forwarded_model_id`` alias is never applied to a different provider. """ async with create_session() as session: query = select(UpstreamProviderRow).options( @@ -158,149 +215,149 @@ async def _load_model_visibility() -> tuple[ ) provider_rows = (await session.exec(query)).all() - overrides_by_id: dict[str, tuple[ModelRow, float]] = {} - disabled_model_ids: set[str] = set() + overrides_by_key: dict[ModelKey, ModelRow] = {} + disabled_model_keys: set[ModelKey] = set() enabled_provider_ids: set[int] = set() for provider in provider_rows: - if not provider.enabled: + if not provider.enabled or provider.id is None: continue - if provider.id is not None: - enabled_provider_ids.add(provider.id) + enabled_provider_ids.add(provider.id) for model in provider.models: + key = (model.id.lower(), provider.id) if model.enabled: - overrides_by_id[model.id] = (model, provider.provider_fee) + overrides_by_key[key] = model else: - disabled_model_ids.add(model.id) + disabled_model_keys.add(key) - return overrides_by_id, disabled_model_ids, enabled_provider_ids - - -def _row_to_visible_model( - model_id: str, - row: ModelRow, - provider_fee: float, -) -> Model | None: - """Convert an enabled DB override row into a routed model object.""" - from ..payment.models import _row_to_model - - try: - return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) - except Exception as exc: # noqa: BLE001 - skip invalid override row - logger.warning( - "Skipping invalid model override while collecting model paths", - extra={ - "model_id": model_id, - "upstream_provider_id": getattr(row, "upstream_provider_id", None), - "error": str(exc), - "error_type": type(exc).__name__, - }, - ) - return None + return overrides_by_key, disabled_model_keys, enabled_provider_ids def _apply_model_visibility( upstream: BaseUpstreamProvider, - overrides_by_id: dict[str, tuple[ModelRow, float]] | None, - disabled_model_ids: set[str] | None, + overrides_by_key: dict[ModelKey, ModelRow] | None, + disabled_model_keys: set[ModelKey] | None, ) -> list[object]: - """Return provider models after DB disabled/override state is applied.""" - overrides_by_id = overrides_by_id or {} - disabled_model_ids = disabled_model_ids or set() + """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. + """ + overrides_by_key = overrides_by_key or {} + disabled_model_keys = disabled_model_keys or set() + upstream_provider_id = getattr(upstream, "db_id", None) + if not isinstance(upstream_provider_id, int): + return [ + model + for model in upstream.get_cached_models() + if getattr(model, "enabled", True) + ] + visible_models: list[object] = [] seen_model_ids: set[str] = set() for model in upstream.get_cached_models(): model_id = getattr(model, "id", "") - if not getattr(model, "enabled", True) or model_id in disabled_model_ids: + key = (model_id.lower(), upstream_provider_id) + if not getattr(model, "enabled", True) or key in disabled_model_keys: continue - - if model_id in overrides_by_id: - override_row, provider_fee = overrides_by_id[model_id] - visible_model = _row_to_visible_model(model_id, override_row, provider_fee) - if visible_model is None: - continue - model = visible_model - - if not getattr(model, "enabled", True): - continue - visible_models.append(model) + # Apply overrides only for this provider's own model row. + override_row = overrides_by_key.get(key) + visible: object = model if override_row is None else override_row + visible_models.append(visible) seen_model_ids.add(model_id.lower()) - upstream_provider_id = getattr(upstream, "db_id", None) - if isinstance(upstream_provider_id, int): - for model_id, (override_row, provider_fee) in overrides_by_id.items(): - if model_id in disabled_model_ids: - continue - if ( - getattr(override_row, "upstream_provider_id", None) - != upstream_provider_id - ): - continue - if model_id.lower() in seen_model_ids: - continue - override_model = _row_to_visible_model(model_id, override_row, provider_fee) - if override_model is None: - continue - if getattr(override_model, "enabled", True): - visible_models.append(override_model) - seen_model_ids.add(model_id.lower()) + # DB-only override rows for this provider with no cached counterpart. + for (model_id_lower, provider_id), override_row in overrides_by_key.items(): + if provider_id != upstream_provider_id: + continue + if model_id_lower in seen_model_ids: + continue + visible_models.append(override_row) + seen_model_ids.add(model_id_lower) return visible_models async def _collect_provider_paths( upstream: BaseUpstreamProvider, - overrides_by_id: dict[str, tuple[ModelRow, float]] | None = None, - disabled_model_ids: set[str] | None = None, -) -> list[tuple[str, str]]: + overrides_by_key: dict[ModelKey, ModelRow] | None = None, + disabled_model_keys: set[ModelKey] | None = None, + cycle: _RefreshCycleState | None = None, +) -> list[tuple[str, str]] | None: """Collect ``(model_id, path)`` pairs for one provider instance. - Emits the direct ```` path for normal upstreams. For - OpenRouter-compatible providers, emits one path per OpenRouter sub-provider - endpoint, prefixed the same way response stamping prefixes it. + Emits the provider's ``discovery_base_paths`` for normal upstreams. For + OpenRouter-compatible providers, additionally emits one path per OpenRouter + sub-provider endpoint via ``discovery_path_for_subprovider`` so the strings + match response stamping exactly. + + Returns ``None`` when the provider's path set could not be determined this + cycle (every endpoint fetch degraded); callers must then keep previously + persisted rows instead of wiping them. """ - provider_type = (upstream.provider_type or "").strip() - models = _apply_model_visibility(upstream, overrides_by_id, disabled_model_ids) - is_openrouter = is_openrouter_base_url(upstream.base_url) + cycle = cycle or _RefreshCycleState() + models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) + base_paths = upstream.discovery_base_paths() - pairs: list[tuple[str, str]] = [] - if not is_openrouter: - for model in models: - if provider_type: - pairs.append((exposed_model_id(model), provider_type)) - return pairs + if not is_openrouter_base_url(upstream.base_url): + return [ + (exposed_model_id(model), path) for model in models for path in base_paths + ] - if not provider_type: - return pairs + if not (upstream.provider_type or "").strip(): + return [] + any_fetch_succeeded = False + any_fetch_attempted = False semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) - async with httpx.AsyncClient() as client: + async with _make_http_client() as client: async def _for_model(model: object) -> list[tuple[str, str]]: + nonlocal any_fetch_succeeded, any_fetch_attempted + model_id = exposed_model_id(model) + # Base paths always apply: responses whose upstream payload lacks a + # provider field are stamped with them (see _apply_provider_field). + pairs = [(model_id, path) for path in base_paths] author_slug = openrouter_author_slug(model) if not author_slug: - return [] - paths = await _fetch_openrouter_endpoint_paths( + return pairs + any_fetch_attempted = True + sub_providers = await _fetch_openrouter_endpoint_subproviders( client, upstream.base_url, upstream.api_key, author_slug, - provider_type, semaphore, + cycle, ) - model_id = exposed_model_id(model) - return [(model_id, path) for path in paths] + if sub_providers is None: + return [] + any_fetch_succeeded = True + paths = [ + upstream.discovery_path_for_subprovider(name) for name in sub_providers + ] + pairs.extend((model_id, path) for path in paths if path) + return list(dict.fromkeys(pairs)) results = await asyncio.gather( *(_for_model(m) for m in models), return_exceptions=True ) + if any_fetch_attempted and not any_fetch_succeeded: + # Every endpoint lookup degraded (offline, throttled, bad payloads): + # the true path set is unknown, not empty. + return None + + pairs: list[tuple[str, str]] = [] for result in results: if isinstance(result, BaseException): logger.warning( "OpenRouter endpoint discovery task errored", - extra={"provider": provider_type, "error": str(result)}, + extra={"provider": upstream.provider_type, "error": str(result)}, ) continue pairs.extend(result) @@ -314,33 +371,57 @@ async def _persist_provider_paths( """Replace all rows for ``upstream_provider_id`` with ``pairs``. Replacement (not upsert) so stale paths disappear when provider config or - upstream availability changes. + upstream availability changes. Rows are written with chunked bulk INSERTs + so the transaction holds SQLite's write lock briefly — billing writes share + this database file. """ unique_pairs = list(dict.fromkeys(pairs)) + now = int(time.time()) async with create_session() as session: await session.exec( # type: ignore[call-overload] delete(ModelPathRow).where( col(ModelPathRow.upstream_provider_id) == upstream_provider_id ) ) - for model_id, path in unique_pairs: - session.add( - ModelPathRow( - model_id=model_id, - path=path, - upstream_provider_id=upstream_provider_id, - ) + for start in range(0, len(unique_pairs), _PERSIST_CHUNK_SIZE): + chunk = unique_pairs[start : start + _PERSIST_CHUNK_SIZE] + await session.execute( + insert(ModelPathRow), + [ + { + "model_id": model_id, + "path": path, + "upstream_provider_id": upstream_provider_id, + "updated_at": now, + } + for model_id, path in chunk + ], ) await session.commit() -async def _prune_inactive_provider_paths(active_provider_ids: set[int]) -> None: - """Delete paths for providers no longer present in the live upstream set.""" +async def prune_model_paths_for_inactive_providers() -> None: + """Delete paths whose provider is no longer enabled in the database. + + Called from ``refresh_model_maps`` so admin mutations (disable/delete + provider) stop advertising a provider's paths immediately instead of + waiting for the next timed refresh. Uses the DB as the source of truth, so + it is safe at boot even before upstreams initialize. + """ async with create_session() as session: + enabled_ids = ( + await session.exec( + select(UpstreamProviderRow.id).where( + col(UpstreamProviderRow.enabled).is_(True) + ) + ) + ).all() stmt = delete(ModelPathRow) - if active_provider_ids: + if enabled_ids: stmt = stmt.where( - col(ModelPathRow.upstream_provider_id).not_in(active_provider_ids) + col(ModelPathRow.upstream_provider_id).not_in( + [pid for pid in enabled_ids if pid is not None] + ) ) await session.exec(stmt) # type: ignore[call-overload] await session.commit() @@ -352,28 +433,42 @@ async def refresh_model_paths( """Recompute and persist model paths for every enabled provider. One provider's failure is logged and isolated; it must not break the rest. + A provider whose paths could not be determined this cycle keeps its + previously persisted rows. An empty ``upstreams`` list (e.g. a failed + ``initialize_upstreams`` at boot) is treated as "unknown" and touches + nothing. """ + if not upstreams: + logger.warning("Skipping model paths refresh: no live upstreams") + return + ( - overrides_by_id, - disabled_model_ids, + overrides_by_key, + disabled_model_keys, enabled_provider_ids, ) = await _load_model_visibility() - active_provider_ids = { - upstream.db_id - for upstream in upstreams - if upstream.db_id is not None and upstream.db_id in enabled_provider_ids - } - await _prune_inactive_provider_paths(active_provider_ids) + await prune_model_paths_for_inactive_providers() + cycle = _RefreshCycleState() for upstream in upstreams: if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: continue try: pairs = await _collect_provider_paths( upstream, - overrides_by_id=overrides_by_id, - disabled_model_ids=disabled_model_ids, + overrides_by_key=overrides_by_key, + disabled_model_keys=disabled_model_keys, + cycle=cycle, ) + if pairs is None: + logger.warning( + "Model paths unknown this cycle; keeping previous rows", + extra={ + "provider": upstream.provider_type or upstream.base_url, + "db_id": upstream.db_id, + }, + ) + continue await _persist_provider_paths(upstream.db_id, pairs) except Exception as e: # noqa: BLE001 - isolate per-provider failures logger.error( @@ -387,18 +482,27 @@ async def refresh_model_paths( ) +def _refresh_interval_seconds() -> int: + """Current interval, re-read every loop so runtime setting changes apply.""" + from ..core.settings import settings + + if not getattr(settings, "enable_model_paths_refresh", True): + return 0 + return int(getattr(settings, "model_paths_refresh_interval_seconds", 0) or 0) + + async def refresh_model_paths_periodically( upstreams_provider: ( Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider] ), ) -> None: - """Background task mirroring ``refresh_upstreams_models_periodically``.""" - from ..core.settings import settings + """Background task mirroring ``refresh_upstreams_models_periodically``. - interval = getattr(settings, "model_paths_refresh_interval_seconds", 0) - if not interval or interval <= 0: - logger.info("Model paths refresh disabled (interval <= 0)") - return + The interval and enable flag are re-read every iteration, so the refresh + can be turned off (or on) and retuned without a restart. While disabled the + task idles instead of exiting, so re-enabling takes effect. + """ + _DISABLED_POLL_SECONDS = 60.0 def _resolve_upstreams() -> list[BaseUpstreamProvider]: if callable(upstreams_provider): @@ -406,6 +510,14 @@ async def refresh_model_paths_periodically( return upstreams_provider while True: + interval = _refresh_interval_seconds() + if interval <= 0: + try: + await asyncio.sleep(_DISABLED_POLL_SECONDS) + except asyncio.CancelledError: + break + continue + try: await refresh_model_paths(_resolve_upstreams()) except asyncio.CancelledError: @@ -423,50 +535,79 @@ async def refresh_model_paths_periodically( break -async def get_all_model_paths() -> list[dict]: +async def get_all_model_paths() -> dict: """All models with their paths, shaped for ``GET /v1/models/paths``.""" async with create_session() as session: rows = ( - await session.exec(select(ModelPathRow).order_by(ModelPathRow.model_id)) + await session.exec( + select(ModelPathRow).order_by( + col(ModelPathRow.model_id), + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) + ) ).all() grouped: dict[str, list[dict]] = {} seen_paths: dict[str, set[str]] = {} + updated_at = 0 for row in rows: + updated_at = max(updated_at, row.updated_at) model_id = public_model_id(row.model_id) if row.path in seen_paths.setdefault(model_id, set()): continue seen_paths[model_id].add(row.path) grouped.setdefault(model_id, []).append({"path": row.path}) - return [{"id": model_id, "paths": paths} for model_id, paths in grouped.items()] + # Deterministic output: models sorted by public id, paths sorted within. + data: list[dict] = [] + for grouped_model_id in sorted(grouped): + model_paths = sorted(grouped[grouped_model_id], key=lambda p: str(p["path"])) + data.append({"id": grouped_model_id, "paths": model_paths}) + return {"data": data, "updated_at": updated_at or None} -async def get_paths_for_model(model_id: str) -> list[dict]: +async def get_paths_for_model(model_id: str) -> dict: """Paths for a single model, shaped for ``GET /v1/models/paths/model``. Match by the public, unqualified model id, mirroring the model cache alias behavior. Both ``deepseek-v4-pro`` and ``deepseek/deepseek-v4-pro`` resolve - every row whose stored id has the same base model id. + every row whose stored id has the same base model id. The candidate set is + narrowed in SQL (exact id or ``%/`` suffix) so the route does not + materialize the whole table per request. """ - requested_id = public_model_id(model_id) + # The request may be a full stored id ("z-ai/glm-5v-turbo") or an + # already-stripped public id ("fireworks/models/glm-5"); accept both. + accepted_ids = {model_id, public_model_id(model_id)} async with create_session() as session: + conditions = [] + for candidate in accepted_ids: + conditions.append(col(ModelPathRow.model_id) == candidate) + conditions.append(col(ModelPathRow.model_id).endswith(f"/{candidate}")) rows = ( await session.exec( - select(ModelPathRow).order_by( - ModelPathRow.path, + select(ModelPathRow) + .where(or_(*conditions)) + .order_by( + col(ModelPathRow.path), col(ModelPathRow.upstream_provider_id), - ModelPathRow.model_id, + col(ModelPathRow.model_id), ) ) ).all() seen: set[str] = set() paths: list[dict] = [] + updated_at = 0 for row in rows: - if public_model_id(row.model_id) != requested_id: + # The SQL suffix match is a prefilter; enforce the exact public-id rule. + if ( + row.model_id not in accepted_ids + and public_model_id(row.model_id) not in accepted_ids + ): continue + updated_at = max(updated_at, row.updated_at) if row.path in seen: continue seen.add(row.path) paths.append({"path": row.path}) - return paths + return {"data": paths, "updated_at": updated_at or None} diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 1caeaa5c..fe92c4f9 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -18,6 +18,23 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): supports_anthropic_messages = True litellm_provider_prefix = "openrouter/" + def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: + """Mirror ``_apply_provider_field``: strip repeated prefixes, map a + missing or self-echoing sub-provider to the literal ``"unknown"``.""" + provider_type = (self.provider_type or "").strip() + sub = (sub_provider or "").strip() + prefix = f"{provider_type}:" + while sub.lower().startswith(prefix.lower()): + sub = sub[len(prefix) :].strip() + if not sub or sub.lower() == provider_type.lower(): + return "unknown" + return f"{provider_type}:{sub}" + + def discovery_base_paths(self) -> list[str]: + """Native OpenRouter never stamps a bare ``openrouter``; a response + with no sub-provider is stamped ``unknown``.""" + return ["unknown"] + def _apply_provider_field(self, response_json: object) -> None: """Stamp the ``provider`` field for OpenRouter responses. diff --git a/tests/unit/test_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index d820a056..c1ce13c8 100644 --- a/tests/unit/test_fee_payout_migration.py +++ b/tests/unit/test_fee_payout_migration.py @@ -37,7 +37,9 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None: "payout_in_progress_msats, payout_started_at FROM routstr_fees" ).fetchone() - assert version == ("9c4d8e2f1a6b",) + # Head of the 7f2843d3f4e4 lineage: model-paths chains onto the fee-payout + # repair migration. + assert version == ("4e0c3d195a49",) assert { "id", "accumulated_msats", diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 947c4312..b1747436 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -1,4 +1,10 @@ -"""Tests for the model-path discovery service and endpoints.""" +"""Tests for the model-path discovery service and endpoints. + +These tests exercise the public entry points (``refresh_model_paths``, +``get_all_model_paths``, ``get_paths_for_model``) rather than private helpers, +and fake OpenRouter at the transport level (``httpx.MockTransport``) so a +signature drift in the SUT fails loudly instead of silently returning ``[]``. +""" from __future__ import annotations @@ -7,11 +13,13 @@ import json import os from contextlib import asynccontextmanager from types import SimpleNamespace -from typing import Any, AsyncGenerator +from typing import Any, AsyncGenerator, Callable +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient +from sqlalchemy import event from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession @@ -22,6 +30,8 @@ os.environ.setdefault("UPSTREAM_API_KEY", "test") from routstr.core.db import ModelRow, UpstreamProviderRow # 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 +from routstr.upstream.openrouter import OpenRouterUpstreamProvider # noqa: E402 # --------------------------------------------------------------------------- # # Fakes @@ -74,7 +84,10 @@ def _model_row( ) -class _FakeProvider: +class _FakeProvider(BaseUpstreamProvider): + """Real ``BaseUpstreamProvider`` so the discovery-path hooks are the + production ones, with cached models injected.""" + def __init__( self, *, @@ -84,37 +97,83 @@ class _FakeProvider: db_id: int | None = 1, api_key: str = "sk-test", ) -> None: - self.provider_type = provider_type - self.base_url = base_url - self.api_key = api_key + super().__init__(base_url=base_url, api_key=api_key) + self.provider_type = provider_type # shadow the class attribute self.db_id = db_id self._models = models - def get_cached_models(self) -> list[SimpleNamespace]: + def get_cached_models(self) -> list[SimpleNamespace]: # type: ignore[override] return self._models -class _FakeResponse: - def __init__(self, status_code: int, payload: Any) -> None: - self.status_code = status_code - self._payload = payload +class _FakeOpenRouterProvider(OpenRouterUpstreamProvider): + """Real OpenRouter provider so the ``unknown`` mapping is the production one.""" - def json(self) -> Any: - return self._payload + def __init__( + self, + *, + models: list[SimpleNamespace], + db_id: int | None = 2, + api_key: str = "sk-or", + ) -> None: + super().__init__(api_key=api_key) + self.db_id = db_id + self._models = models + + def get_cached_models(self) -> list[SimpleNamespace]: # type: ignore[override] + return self._models + + +def _mock_transport( + monkeypatch: pytest.MonkeyPatch, + handler: Callable[[httpx.Request], httpx.Response], +) -> dict[str, int]: + """Route the SUT's HTTP through ``httpx.MockTransport`` and count requests.""" + counter = {"requests": 0} + + def _counting_handler(request: httpx.Request) -> httpx.Response: + counter["requests"] += 1 + return handler(request) + + def _factory() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(_counting_handler)) + + monkeypatch.setattr(mp, "_make_http_client", _factory) + return counter + + +def _endpoints_response(*provider_names: str) -> httpx.Response: + return httpx.Response( + 200, + json={"data": {"endpoints": [{"provider_name": n} for n in provider_names]}}, + ) + + +_SEEDED_PROVIDER_IDS = (1, 2, 4, 5, 7) @pytest.fixture async def patched_session( monkeypatch: pytest.MonkeyPatch, ) -> AsyncGenerator[AsyncEngine, None]: - """Bind the service's ``create_session`` to a fresh in-memory engine.""" + """Bind the service's ``create_session`` to a fresh in-memory engine. + + Foreign keys are enforced (``PRAGMA foreign_keys=ON``) so a ModelPathRow + insert for an unseeded provider fails here even though production SQLite + currently runs with the pragma off. + """ engine = create_async_engine("sqlite+aiosqlite:///:memory:") + + @event.listens_for(engine.sync_engine, "connect") + def _enable_fk(dbapi_conn: Any, _record: Any) -> None: + dbapi_conn.execute("PRAGMA foreign_keys=ON") + async with engine.begin() as conn: await conn.run_sync(SQLModel.metadata.create_all) - # Seed the FK target so ModelPathRow inserts satisfy the constraint. + # Seed every provider id the tests insert path rows for. async with AsyncSession(engine) as session: - for pid in (1, 2): + for pid in _SEEDED_PROVIDER_IDS: session.add( UpstreamProviderRow( id=pid, @@ -136,6 +195,17 @@ async def patched_session( await engine.dispose() +def _paths_of(payload: dict, model_id: str) -> set[str]: + for entry in payload["data"]: + if entry["id"] == model_id: + return {p["path"] for p in entry["paths"]} + return set() + + +def _ids_of(payload: dict) -> set[str]: + return {entry["id"] for entry in payload["data"]} + + # --------------------------------------------------------------------------- # # Predicates / pure helpers # --------------------------------------------------------------------------- # @@ -161,12 +231,18 @@ def test_exposed_model_id_prefers_forwarded() -> None: assert mp.exposed_model_id(_model("claude-x")) == "claude-x" -def test_public_model_id_strips_provider_prefix() -> None: +def test_public_model_id_strips_first_provider_prefix() -> None: + """Must match ``create_model_mappings.get_base_model_id`` (first slash), + so the id shown by discovery can be sent to chat completions verbatim.""" assert mp.public_model_id("z-ai/glm-5v-turbo") == "glm-5v-turbo" assert mp.public_model_id("gpt-4o-mini") == "gpt-4o-mini" + assert ( + mp.public_model_id("accounts/fireworks/models/glm-5") + == "fireworks/models/glm-5" + ) -def test_openrouter_author_slug_uses_canonical_not_forwarded() -> None: +def test_openrouter_author_slug_prefers_canonical() -> None: m = _model( "claude-opus-4.6", forwarded_model_id="forwarded-only", @@ -180,216 +256,174 @@ def test_openrouter_author_slug_falls_back_to_slash_id() -> None: assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" +def test_openrouter_author_slug_falls_back_to_forwarded_id() -> None: + """Admin-created alias rows have a slash-less local id; the forwarded id is + what the proxy actually sends to OpenRouter, so it is a usable slug.""" + m = _model("my-alias", forwarded_model_id="anthropic/claude-opus-4.6") + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + def test_openrouter_author_slug_none_when_no_slash() -> None: m = _model("claude-opus-4.6", canonical_slug="claude-opus-4.6") assert mp.openrouter_author_slug(m) is None +def test_discovery_paths_mirror_response_stamping() -> None: + """The discovery hook and ``_apply_provider_field`` must agree.""" + generic = _FakeProvider(provider_type="generic", base_url="https://x", models=[]) + assert generic.discovery_path_for_subprovider("Anthropic") == "generic:Anthropic" + assert generic.discovery_base_paths() == ["generic"] + + native = _FakeOpenRouterProvider(models=[]) + assert native.discovery_path_for_subprovider("GMICloud") == "openrouter:GMICloud" + # Sub-provider echoing the router name is stamped "unknown" on responses. + assert native.discovery_path_for_subprovider("OpenRouter") == "unknown" + assert native.discovery_path_for_subprovider("openrouter:openrouter") == "unknown" + assert native.discovery_path_for_subprovider(None) == "unknown" + assert native.discovery_base_paths() == ["unknown"] + + # --------------------------------------------------------------------------- # -# Collection +# Refresh through the public entry point # --------------------------------------------------------------------------- # @pytest.mark.asyncio -async def test_direct_provider_single_path_uses_provider_type() -> None: +async def test_direct_provider_single_path_uses_provider_type( + patched_session: AsyncEngine, +) -> None: provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", models=[_model("claude-opus-4.6")], + db_id=1, ) - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("claude-opus-4.6", "anthropic")] + await mp.refresh_model_paths([provider]) + payload = await mp.get_all_model_paths() + assert payload["data"] == [ + {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + ] + assert payload["updated_at"] is not None @pytest.mark.asyncio -async def test_direct_path_stores_exposed_model_id() -> None: +async def test_direct_path_stores_exposed_model_id( + patched_session: AsyncEngine, +) -> None: provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", models=[_model("internal-id", forwarded_model_id="claude-opus-4.6")], + db_id=1, ) - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("claude-opus-4.6", "anthropic")] + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} @pytest.mark.asyncio -async def test_disabled_models_excluded() -> None: +async def test_disabled_cached_models_excluded( + patched_session: AsyncEngine, +) -> None: provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", - models=[ - _model("enabled-model"), - _model("disabled-model", enabled=False), - ], + models=[_model("enabled-model"), _model("disabled-model", enabled=False)], + db_id=1, ) - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("enabled-model", "anthropic")] + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"enabled-model"} @pytest.mark.asyncio -async def test_openrouter_provider_adds_endpoint_paths( - monkeypatch: pytest.MonkeyPatch, -) -> None: - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _fake_get( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - return _FakeResponse( - 200, - { - "data": { - "endpoints": [ - {"provider_name": "Anthropic"}, - {"provider_name": "Amazon Bedrock"}, - ] - } - }, - ) - - monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert ("claude-opus-4.6", "openrouter:Anthropic") in pairs - assert ("claude-opus-4.6", "openrouter:Amazon Bedrock") in pairs - assert ("claude-opus-4.6", "openrouter") not in pairs - assert len(pairs) == 2 - - -@pytest.mark.asyncio -async def test_generic_provider_with_openrouter_base_url_discovers( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """A generic provider pointed at OpenRouter exposes the response-stamped - ``generic:`` path, not a native ``openrouter:`` path.""" - provider = _FakeProvider( - provider_type="generic", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _fake_get( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - return _FakeResponse( - 200, {"data": {"endpoints": [{"provider_name": "Anthropic"}]}} - ) - - monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("claude-opus-4.6", "generic:Anthropic")] - - -@pytest.mark.asyncio -async def test_openrouter_failure_degrades_gracefully( - monkeypatch: pytest.MonkeyPatch, -) -> None: - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _boom( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - raise RuntimeError("network down") - - monkeypatch.setattr("httpx.AsyncClient.get", _boom) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [] - - -@pytest.mark.asyncio -async def test_openrouter_rate_limit_skips_model( - monkeypatch: pytest.MonkeyPatch, -) -> None: - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _rate_limited( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - return _FakeResponse(429, {}) - - monkeypatch.setattr("httpx.AsyncClient.get", _rate_limited) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [] - - -@pytest.mark.asyncio -async def test_openrouter_fanout_is_bounded( - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) - models = [ - _model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20) - ] - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=models, - ) - - state = {"current": 0, "max": 0} - - async def _slow_get( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - state["current"] += 1 - state["max"] = max(state["max"], state["current"]) - await asyncio.sleep(0.02) - state["current"] -= 1 - return _FakeResponse(200, {"data": {"endpoints": [{"provider_name": "X"}]}}) - - monkeypatch.setattr("httpx.AsyncClient.get", _slow_get) - - await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" - - -# --------------------------------------------------------------------------- # -# Persistence + query -# --------------------------------------------------------------------------- # - - -@pytest.mark.asyncio -async def test_refresh_replaces_stale_rows( +async def test_disabling_model_on_one_provider_keeps_other_provider( patched_session: AsyncEngine, ) -> None: - await mp._persist_provider_paths(1, [("m1", "anthropic"), ("m2", "anthropic")]) - first = await mp.get_all_model_paths() - assert {row["id"] for row in first} == {"m1", "m2"} + """Regression for cross-provider isolation: ModelRow's primary key is + (id, upstream_provider_id), so a disable row on provider 2 must not hide + provider 1's model.""" + async with AsyncSession(patched_session) as session: + session.add(_model_row("shared-model", upstream_provider_id=2, enabled=False)) + await session.commit() - # Second refresh with a different set — stale m2 must disappear. - await mp._persist_provider_paths(1, [("m1", "anthropic")]) - second = await mp.get_all_model_paths() - assert {row["id"] for row in second} == {"m1"} + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("shared-model")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("shared-model")], + db_id=2, + ) + + await mp.refresh_model_paths([p1, p2]) + + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "shared-model") == {"anthropic"} + + +@pytest.mark.asyncio +async def test_override_alias_not_applied_across_providers( + patched_session: AsyncEngine, +) -> None: + """Provider 2's forwarded_model_id must never rename provider 1's model.""" + async with AsyncSession(patched_session) as session: + session.add( + _model_row( + "shared-model", + upstream_provider_id=2, + forwarded_model_id="private-alias", + ) + ) + await session.commit() + + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("shared-model")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("shared-model")], + db_id=2, + ) + + await mp.refresh_model_paths([p1, p2]) + + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "shared-model") == {"anthropic"} + assert _paths_of(payload, "private-alias") == {"generic"} + + +@pytest.mark.asyncio +async def test_override_matching_is_case_insensitive( + patched_session: AsyncEngine, +) -> None: + """Routing lowercases both sides when matching DB rows to cached models; + discovery must do the same for mixed-case ids.""" + async with AsyncSession(patched_session) as session: + session.add( + _model_row( + "deepseek-ai/deepseek-v4-flash", + upstream_provider_id=1, + forwarded_model_id="public-alias", + ) + ) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("deepseek-ai/DeepSeek-V4-Flash")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"public-alias"} @pytest.mark.asyncio @@ -407,9 +441,9 @@ async def test_refresh_model_paths_excludes_db_disabled_override( db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [] + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio @@ -427,9 +461,9 @@ async def test_refresh_model_paths_uses_db_forwarded_alias( db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [ + assert (await mp.get_all_model_paths())["data"] == [ {"id": "public-alias", "paths": [{"path": "anthropic"}]} ] @@ -449,44 +483,88 @@ async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cac db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [ + assert (await mp.get_all_model_paths())["data"] == [ {"id": "public-deployment", "paths": [{"path": "generic"}]} ] @pytest.mark.asyncio -async def test_refresh_model_paths_prunes_inactive_provider_rows( +async def test_refresh_replaces_stale_rows( patched_session: AsyncEngine, ) -> None: - await mp._persist_provider_paths(1, [("m1", "anthropic")]) - await mp._persist_provider_paths(2, [("m2", "openrouter:Anthropic")]) + p_two_models = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1"), _model("m2")], + db_id=1, + ) + await mp.refresh_model_paths([p_two_models]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1", "m2"} + p_one_model = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([p_one_model]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + +@pytest.mark.asyncio +async def test_refresh_with_no_upstreams_keeps_existing_rows( + patched_session: AsyncEngine, +) -> None: + """An empty live upstream list (e.g. failed boot init) means "unknown", + not "delete everything".""" provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", models=[_model("m1")], db_id=1, ) - - await mp.refresh_model_paths([provider]) # type: ignore[list-item] - active_only = await mp.get_all_model_paths() - assert active_only == [{"id": "m1", "paths": [{"path": "anthropic"}]}] + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} await mp.refresh_model_paths([]) - assert await mp.get_all_model_paths() == [] + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + +@pytest.mark.asyncio +async def test_prune_removes_rows_of_disabled_db_provider( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.enabled = False + session.add(provider_row) + await session.commit() + + await mp.prune_model_paths_for_inactive_providers() + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio async def test_refresh_model_paths_skips_disabled_db_provider( patched_session: AsyncEngine, ) -> None: - await mp._persist_provider_paths(1, [("stale-model", "anthropic")]) async with AsyncSession(patched_session) as session: provider_row = await session.get(UpstreamProviderRow, 1) assert provider_row is not None provider_row.enabled = False + session.add(provider_row) await session.commit() provider = _FakeProvider( @@ -496,109 +574,9 @@ async def test_refresh_model_paths_skips_disabled_db_provider( db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [] - - -@pytest.mark.asyncio -async def test_same_model_two_providers_two_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) - - data = await mp.get_all_model_paths() - assert len(data) == 1 - entry = data[0] - assert entry["id"] == "claude-opus-4.6" - paths = {p["path"] for p in entry["paths"]} - assert paths == {"anthropic", "openrouter:Anthropic"} - # No canonical_id anywhere. - assert "canonical_id" not in entry - assert all("canonical_id" not in p for p in entry["paths"]) - - -@pytest.mark.asyncio -async def test_get_all_model_paths_deduplicates_visible_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) - - assert await mp.get_all_model_paths() == [ - {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} - ] - - -@pytest.mark.asyncio -async def test_get_all_model_paths_returns_unqualified_model_ids( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) - await mp._persist_provider_paths(5, [("openai/gpt-4o-mini", "openrouter:OpenAI")]) - - data = await mp.get_all_model_paths() - - assert {row["id"] for row in data} == {"glm-5v-turbo", "gpt-4o-mini"} - - -@pytest.mark.asyncio -async def test_get_paths_for_model_returns_only_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) - - paths = await mp.get_paths_for_model("claude-opus-4.6") - assert {p["path"] for p in paths} == {"anthropic", "openrouter:Anthropic"} - assert all(set(p.keys()) == {"path"} for p in paths) - assert await mp.get_paths_for_model("does-not-exist") == [] - - -@pytest.mark.asyncio -async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) - - paths = await mp.get_paths_for_model("glm-5v-turbo") - - assert paths == [{"path": "openrouter:Z.AI"}] - - -@pytest.mark.asyncio -async def test_get_paths_for_model_deduplicates_visible_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) - - assert await mp.get_paths_for_model("claude-opus-4.6") == [ - {"path": "anthropic"} - ] - assert await mp.get_paths_for_model("anthropic/claude-opus-4.6") == [ - {"path": "anthropic"} - ] - - -@pytest.mark.asyncio -async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(7, [("deepseek-v4-pro", "generic")]) - await mp._persist_provider_paths( - 4, [("deepseek/deepseek-v4-pro", "openrouter:DeepSeek")] - ) - - short_paths = await mp.get_paths_for_model("deepseek-v4-pro") - prefixed_paths = await mp.get_paths_for_model("deepseek/deepseek-v4-pro") - - assert short_paths == [ - {"path": "generic"}, - {"path": "openrouter:DeepSeek"}, - ] - assert prefixed_paths == short_paths + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio @@ -611,8 +589,8 @@ async def test_refresh_model_paths_skips_provider_without_db_id( models=[_model("claude-opus-4.6")], db_id=None, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] - assert await mp.get_all_model_paths() == [] + await mp.refresh_model_paths([provider]) + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio @@ -626,30 +604,476 @@ async def test_refresh_model_paths_isolates_provider_failure( db_id=1, ) bad = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("m", canonical_slug="a/m")], + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("m")], db_id=2, ) original = mp._collect_provider_paths - async def _maybe_fail( - upstream: Any, *args: Any, **kwargs: Any - ) -> list[tuple[str, str]]: + async def _maybe_fail(upstream: Any, *args: Any, **kwargs: Any) -> Any: if upstream is bad: raise RuntimeError("boom") - return await original(upstream, *args, **kwargs) # type: ignore[arg-type] + return await original(upstream, *args, **kwargs) monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) - await mp.refresh_model_paths([good, bad]) # type: ignore[list-item] - data = await mp.get_all_model_paths() - assert {row["id"] for row in data} == {"claude-opus-4.6"} + await mp.refresh_model_paths([good, bad]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} # --------------------------------------------------------------------------- # -# Endpoints +# OpenRouter endpoint discovery (transport-level fakes) +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_openrouter_provider_adds_endpoint_paths( + 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, + ) + _mock_transport( + monkeypatch, + lambda request: _endpoints_response("Anthropic", "Amazon Bedrock"), + ) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:Anthropic" in paths + assert "openrouter:Amazon Bedrock" in paths + assert "openrouter" not in paths + + +@pytest.mark.asyncio +async def test_openrouter_self_echoing_subprovider_maps_to_unknown( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Responses stamped for a sub-provider echoing "OpenRouter" say + ``unknown``; discovery must advertise the same string, never + ``openrouter:OpenRouter``.""" + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("OpenRouter")) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:OpenRouter" not in paths + assert "unknown" in paths + + +@pytest.mark.asyncio +async def test_generic_provider_with_openrouter_base_url_discovers( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A generic provider pointed at OpenRouter exposes both the bare + ``generic`` path (stamped when the upstream omits its provider field) and + the ``generic:`` endpoint paths.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=1, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert paths == {"generic", "generic:Anthropic"} + + +@pytest.mark.asyncio +async def test_openrouter_failure_keeps_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A transient upstream failure means the path set is unknown; previously + persisted rows must survive, mirroring ``refresh_models_cache``.""" + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:Anthropic" in before + + def _network_down(request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("network down", request=request) + + _mock_transport(monkeypatch, _network_down) + await mp.refresh_model_paths([provider]) + + after = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert after == before + + +@pytest.mark.asyncio +async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """The first 429 latches: no further endpoint requests this cycle, and the + provider's previously persisted rows survive.""" + models = [_model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(10)] + provider = _FakeOpenRouterProvider(models=models, db_id=2) + + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + assert _paths_of(await mp.get_all_model_paths(), "m0") == { + "unknown", + "openrouter:Anthropic", + } + + counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) + await mp.refresh_model_paths([provider]) + + # Up to _OPENROUTER_CONCURRENCY requests may already be in flight when the + # first 429 lands; the latch must stop everything after that. + assert counter["requests"] <= mp._OPENROUTER_CONCURRENCY, ( + "429 must abort the remaining fan-out" + ) + assert _paths_of(await mp.get_all_model_paths(), "m0") == { + "unknown", + "openrouter:Anthropic", + } + + +@pytest.mark.asyncio +async def test_openrouter_bad_payload_shapes_do_not_raise( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """``endpoints: null`` and non-list endpoint payloads are swallowed as + documented, not raised into the generic task-errored bucket.""" + for payload in ( + {"data": {"endpoints": None}}, + {"data": {"endpoints": "none"}}, + {"data": None}, + {}, + ): + provider = _FakeOpenRouterProvider( + models=[_model("m", canonical_slug="a/m")], db_id=2 + ) + _mock_transport( + monkeypatch, lambda request, p=payload: httpx.Response(200, json=p) + ) + # Must not raise. + await mp.refresh_model_paths([provider]) + + +@pytest.mark.asyncio +async def test_openrouter_shared_base_url_fetched_once( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Two providers on the same OpenRouter base URL share the per-cycle + endpoint cache instead of fetching byte-identical bodies twice.""" + native = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + generic = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=4, + ) + counter = _mock_transport( + monkeypatch, lambda request: _endpoints_response("Anthropic") + ) + + await mp.refresh_model_paths([native, generic]) + + assert counter["requests"] == 1 + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:Anthropic" in paths + assert "generic:Anthropic" in paths + + +@pytest.mark.asyncio +async def test_openrouter_fanout_is_bounded( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) + models = [_model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20)] + provider = _FakeOpenRouterProvider(models=models, db_id=2) + + state = {"current": 0, "max": 0} + + async def _slow_handler(request: httpx.Request) -> httpx.Response: + state["current"] += 1 + state["max"] = max(state["max"], state["current"]) + await asyncio.sleep(0.02) + state["current"] -= 1 + return _endpoints_response("X") + + def _factory() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(_slow_handler)) + + monkeypatch.setattr(mp, "_make_http_client", _factory) + + await mp.refresh_model_paths([provider]) + assert state["max"] > 0, "transport fake was never exercised" + assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" + + +# --------------------------------------------------------------------------- # +# Query endpoints +# --------------------------------------------------------------------------- # + + +async def _seed_two_provider_shared_model(engine: AsyncEngine) -> None: + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other/v1", + models=[_model("claude-opus-4.6")], + db_id=2, + ) + await mp.refresh_model_paths([p1, p2]) + + +@pytest.mark.asyncio +async def test_same_model_two_providers_two_paths( + patched_session: AsyncEngine, +) -> None: + await _seed_two_provider_shared_model(patched_session) + + payload = await mp.get_all_model_paths() + assert len(payload["data"]) == 1 + entry = payload["data"][0] + assert entry["id"] == "claude-opus-4.6" + assert {p["path"] for p in entry["paths"]} == {"anthropic", "generic"} + assert "canonical_id" not in entry + assert all("canonical_id" not in p for p in entry["paths"]) + + +@pytest.mark.asyncio +async def test_get_all_model_paths_deduplicates_visible_paths( + patched_session: AsyncEngine, +) -> None: + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("anthropic/claude-opus-4.6")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=2, + ) + await mp.refresh_model_paths([p1, p2]) + + assert (await mp.get_all_model_paths())["data"] == [ + {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + ] + + +@pytest.mark.asyncio +async def test_get_all_model_paths_is_deterministic( + patched_session: AsyncEngine, +) -> None: + """Output must not depend on rowid insertion order, which changes every + refresh cycle.""" + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("b-model"), _model("a-model")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + first = await mp.get_all_model_paths() + await mp.refresh_model_paths([provider]) + second = await mp.get_all_model_paths() + assert first["data"] == second["data"] + assert [e["id"] for e in first["data"]] == ["a-model", "b-model"] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_returns_only_paths( + patched_session: AsyncEngine, +) -> None: + await _seed_two_provider_shared_model(patched_session) + + payload = await mp.get_paths_for_model("claude-opus-4.6") + assert {p["path"] for p in payload["data"]} == {"anthropic", "generic"} + assert all(set(p.keys()) == {"path"} for p in payload["data"]) + assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("z-ai/glm-5v-turbo")], + db_id=4, + ) + await mp.refresh_model_paths([provider]) + + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ + {"path": "generic"} + ] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( + patched_session: AsyncEngine, +) -> None: + p1 = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("deepseek-v4-pro")], + db_id=7, + ) + p2 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("deepseek/deepseek-v4-pro")], + db_id=4, + ) + await mp.refresh_model_paths([p1, p2]) + + short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] + prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] + + assert {p["path"] for p in short_paths} == {"generic", "anthropic"} + assert prefixed_paths == short_paths + + +@pytest.mark.asyncio +async def test_get_paths_for_model_multi_segment_id_matches_models_listing( + patched_session: AsyncEngine, +) -> None: + """For three-segment ids the discovery id must be the same base id the + rest of the system exposes (first-slash rule), not the last segment.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("accounts/fireworks/models/glm-5")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} + assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ + {"path": "generic"} + ] + assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ + "data" + ] == [{"path": "generic"}] + + +# --------------------------------------------------------------------------- # +# Periodic refresh loop +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_refresh_loop_rereads_interval_and_picks_up_providers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 1, raising=False + ) + + seen_batches: list[list[Any]] = [] + + async def _fake_refresh(upstreams: list[Any]) -> None: + seen_batches.append(list(upstreams)) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + sleeps: list[float] = [] + + async def _fast_sleep(seconds: float) -> None: + sleeps.append(seconds) + if len(seen_batches) >= 2: + raise asyncio.CancelledError + + monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) + + batches = [["p1"], ["p1", "p2"]] + + def _provider() -> list[Any]: + return batches[min(len(seen_batches), len(batches) - 1)] + + await mp.refresh_model_paths_periodically(_provider) # type: ignore[arg-type] + + assert seen_batches[0] == ["p1"] + assert seen_batches[1] == ["p1", "p2"], "loop must re-resolve upstreams each cycle" + assert all(s >= 1 for s in sleeps) + + +@pytest.mark.asyncio +async def test_refresh_loop_idles_while_disabled( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A non-positive interval (or the kill switch) must idle the loop, not + exit it, so runtime re-enabling takes effect.""" + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + + refresh_calls: list[Any] = [] + + async def _fake_refresh(upstreams: list[Any]) -> None: + refresh_calls.append(upstreams) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + idle_sleeps: list[float] = [] + + async def _fast_sleep(seconds: float) -> None: + idle_sleeps.append(seconds) + if len(idle_sleeps) >= 2: + raise asyncio.CancelledError + + monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) + + await mp.refresh_model_paths_periodically(lambda: [object()]) + + assert refresh_calls == [], "disabled loop must not refresh" + assert len(idle_sleeps) == 2, "disabled loop must keep polling, not exit" + + +def test_refresh_interval_respects_kill_switch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + assert mp._refresh_interval_seconds() == 0 + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + assert mp._refresh_interval_seconds() == 600 + + +# --------------------------------------------------------------------------- # +# HTTP endpoints # --------------------------------------------------------------------------- # @@ -662,16 +1086,19 @@ def _make_model_paths_app() -> FastAPI: def test_model_paths_endpoint_returns_all_paths( monkeypatch: pytest.MonkeyPatch, ) -> None: - async def _fake_get_all_model_paths() -> list[dict[str, Any]]: - return [ - { - "id": "claude-opus-4.6", - "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"}, - ], - } - ] + async def _fake_get_all_model_paths() -> dict[str, Any]: + return { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"}, + ], + } + ], + "updated_at": 1753500000, + } monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) @@ -687,7 +1114,8 @@ def test_model_paths_endpoint_returns_all_paths( {"path": "openrouter:Anthropic"}, ], } - ] + ], + "updated_at": 1753500000, } @@ -696,9 +1124,9 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) -> None: calls: list[str] = [] - async def _fake_get_paths_for_model(model_id: str) -> list[dict[str, Any]]: + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: calls.append(model_id) - return [{"path": "generic:Anthropic"}] + return {"data": [{"path": "generic:Anthropic"}], "updated_at": None} monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) @@ -708,5 +1136,8 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) assert response.status_code == 200 - assert response.json() == {"data": [{"path": "generic:Anthropic"}]} + assert response.json() == { + "data": [{"path": "generic:Anthropic"}], + "updated_at": None, + } assert calls == ["anthropic/claude-opus-4.6"] From 06dba681c53250e038aa8434517df00eef6c7802 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:26:49 +0200 Subject: [PATCH 07/14] fix: satisfy strict mypy in model-paths tests Replace untyped lambdas with typed handler/provider functions; CI runs mypy over tests as well. --- tests/unit/test_model_paths.py | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index b1747436..7d3393fe 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -13,7 +13,7 @@ import json import os from contextlib import asynccontextmanager from types import SimpleNamespace -from typing import Any, AsyncGenerator, Callable +from typing import Any, AsyncGenerator, Callable, cast import httpx import pytest @@ -760,9 +760,13 @@ async def test_openrouter_bad_payload_shapes_do_not_raise( provider = _FakeOpenRouterProvider( models=[_model("m", canonical_slug="a/m")], db_id=2 ) - _mock_transport( - monkeypatch, lambda request, p=payload: httpx.Response(200, json=p) - ) + + def _handler( + request: httpx.Request, p: dict[str, Any] | None = payload + ) -> httpx.Response: + return httpx.Response(200, json=p) + + _mock_transport(monkeypatch, _handler) # Must not raise. await mp.refresh_model_paths([provider]) @@ -1051,7 +1055,10 @@ async def test_refresh_loop_idles_while_disabled( monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) - await mp.refresh_model_paths_periodically(lambda: [object()]) + def _upstreams() -> list[BaseUpstreamProvider]: + return [cast(BaseUpstreamProvider, object())] + + await mp.refresh_model_paths_periodically(_upstreams) assert refresh_calls == [], "disabled loop must not refresh" assert len(idle_sleeps) == 2, "disabled loop must keep polling, not exit" From c5da73f1e97bbd670c48ea48a5c81fee27b1677e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 20:12:12 +0200 Subject: [PATCH 08/14] revert: remove unrelated repository formatting --- routstr/algorithm.py | 20 +- routstr/balance.py | 29 +-- routstr/core/admin.py | 26 +- routstr/core/log_manager.py | 21 +- routstr/core/usage_analytics_store.py | 29 +-- routstr/nostr/analytics.py | 8 +- routstr/payment/cost_calculation.py | 19 +- routstr/payment/usage.py | 4 +- routstr/upstream/azure.py | 4 +- routstr/upstream/ehbp.py | 32 ++- routstr/upstream/gemini.py | 4 +- routstr/upstream/gemini_messages.py | 4 +- routstr/upstream/groq.py | 4 +- routstr/upstream/litellm_routing.py | 4 +- routstr/upstream/messages_dispatch.py | 12 +- routstr/upstream/model_paths.py | 322 +++++++++++++++---------- routstr/upstream/ollama.py | 8 +- routstr/upstream/rate_limit.py | 4 +- routstr/upstream/request_correction.py | 4 +- routstr/upstream/routstr.py | 3 +- routstr/upstream/xai.py | 4 +- routstr/wallet.py | 13 +- 22 files changed, 319 insertions(+), 259 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index ef4b8574..fbc5388e 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -232,10 +232,7 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if ( - model_to_use.forwarded_model_id - and model_to_use.forwarded_model_id not in aliases - ): + if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: aliases.append(model_to_use.forwarded_model_id) # Try to set each alias @@ -325,10 +322,7 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if ( - model_to_use.forwarded_model_id - and model_to_use.forwarded_model_id not in aliases - ): + if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: aliases.append(model_to_use.forwarded_model_id) for alias in aliases: @@ -348,10 +342,16 @@ def create_model_mappings( forwarded_model_ids, the one whose forwarded_model_id equals the requested alias wins. """ - if model.forwarded_model_id and model.forwarded_model_id.lower() == alias: + if ( + model.forwarded_model_id + and model.forwarded_model_id.lower() == alias + ): return 5 - if model.id and model.id.lower() == alias: + if ( + model.id + and model.id.lower() == alias + ): return 4 model_base = get_base_model_id(model.id) diff --git a/routstr/balance.py b/routstr/balance.py index 03dc4d33..91b19ce5 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -260,11 +260,7 @@ async def _lookup_key_no_create( async def _restore_balance( - session: AsyncSession, - hashed_key: str, - balance: int, - reserved_balance: int, - mint_url: str, + session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -279,11 +275,7 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={ - "hashed_key": hashed_key, - "restored_balance": balance, - "mint_url": mint_url, - }, + extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, ) @@ -468,23 +460,11 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) + await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) + await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") error_msg = str(e) logger.error( "refund_wallet_endpoint: mint/send failed", @@ -705,6 +685,7 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} + @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 1521510e..66a1d288 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -68,9 +68,7 @@ async def require_admin_api(request: Request) -> None: async with create_session() as session: result = await session.exec(select(CliToken).where(CliToken.token == token)) cli_token = result.first() - if cli_token and ( - cli_token.expires_at is None or cli_token.expires_at > now_ts - ): + if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts): cli_token.last_used_at = now_ts session.add(cli_token) await session.commit() @@ -257,12 +255,16 @@ async def update_password(request: Request, password_update: PasswordUpdate) -> secret = await get_secret(session) if not secret.admin_password_hash: - raise HTTPException(status_code=500, detail="Admin password not configured") + raise HTTPException( + status_code=500, detail="Admin password not configured" + ) if not vault.verify_password( password_update.current_password, secret.admin_password_hash ): - raise HTTPException(status_code=401, detail="Current password is incorrect") + raise HTTPException( + status_code=401, detail="Current password is incorrect" + ) # Validate new password new_password = password_update.new_password.strip() @@ -978,7 +980,9 @@ async def update_upstream_provider_by_slug( lookup = _validate_slug(payload.slug) async with create_session() as session: result = await session.exec( - select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup) + select(UpstreamProviderRow).where( + UpstreamProviderRow.slug == lookup + ) ) provider = result.first() if not provider: @@ -1665,11 +1669,7 @@ async def get_transactions_api( ) total = count_result.one() - stmt = ( - base.order_by(col(CashuTransaction.created_at).desc()) - .offset(offset) - .limit(limit) - ) + stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit) results = await session.exec(stmt) transactions = results.all() @@ -1679,7 +1679,9 @@ async def get_transactions_api( } -@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)]) +@admin_router.get( + "/api/lightning-invoices", dependencies=[Depends(require_admin_api)] +) async def get_lightning_invoices_api( status: str | None = None, purpose: str | None = None, diff --git a/routstr/core/log_manager.py b/routstr/core/log_manager.py index b111f68a..0444dcbf 100644 --- a/routstr/core/log_manager.py +++ b/routstr/core/log_manager.py @@ -408,9 +408,7 @@ class LogManager: def get_error_details(self, hours: int = 24, limit: int = 100) -> dict: def compute() -> dict: try: - return self._usage_store.get_error_details( - hours_back=hours, limit=limit - ) + return self._usage_store.get_error_details(hours_back=hours, limit=limit) except Exception as e: logger.error( f"Usage analytics index failed, falling back to log scan: {e}" @@ -630,7 +628,8 @@ class LogManager: stats["total_tokens"] += input_tokens + output_tokens failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: stats["total_requests"] += 1 @@ -788,9 +787,7 @@ class LogManager: if bucket_key: model_mix_buckets[bucket_key][model] += 1 if revenue_msats > 0: - model_mix_revenue_buckets[bucket_key][model] += ( - revenue_msats - ) + model_mix_revenue_buckets[bucket_key][model] += revenue_msats model_mix_revenue_totals[model] += revenue_msats if input_tokens > 0 or output_tokens > 0: token_total = input_tokens + output_tokens @@ -804,7 +801,8 @@ class LogManager: bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: summary_stats["total_requests"] += 1 @@ -874,7 +872,9 @@ class LogManager: models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True) latest_errors = [ item - for _, item in sorted(latest_errors_heap, key=lambda x: x[0], reverse=True) + for _, item in sorted( + latest_errors_heap, key=lambda x: x[0], reverse=True + ) ] top_model_limit = max(1, min(model_limit, 20)) top_models_requests = [ @@ -1051,7 +1051,8 @@ class LogManager: bucket["warnings"] += 1 failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: bucket["total_requests"] += 1 diff --git a/routstr/core/usage_analytics_store.py b/routstr/core/usage_analytics_store.py index 36fa4bcb..7ba90e24 100644 --- a/routstr/core/usage_analytics_store.py +++ b/routstr/core/usage_analytics_store.py @@ -314,7 +314,9 @@ class UsageAnalyticsStore: if column in existing_columns: return - conn.execute(f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}") + conn.execute( + f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}" + ) logger.info(f"Migrated analytics schema: added {table}.{column}") def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None: @@ -362,11 +364,7 @@ class UsageAnalyticsStore: self._drop_index_tables_locked(conn) self._initialize_schema_locked(conn) - files = ( - log_files - if log_files is not None - else sorted(self.logs_dir.glob("app_*.log")) - ) + files = log_files if log_files is not None else sorted(self.logs_dir.glob("app_*.log")) for log_file in files: try: self._process_log_file_locked(conn, log_file, force_full_read=True) @@ -570,7 +568,8 @@ class UsageAnalyticsStore: model_bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: bucket["total_requests"] += 1 @@ -593,9 +592,9 @@ class UsageAnalyticsStore: if isinstance(max_cost, (int, float)) and max_cost > 0: max_cost_float = float(max_cost) bucket["refunds_msats"] += max_cost_float - model_updates[(minute_key, model)]["refunds_msats"] += ( - max_cost_float - ) + model_updates[(minute_key, model)][ + "refunds_msats" + ] += max_cost_float return ( end_offset, @@ -1033,9 +1032,7 @@ class UsageAnalyticsStore: """, (cutoff_timestamp,), ).fetchone() - total_error_count = ( - int(total_error_count_row[0]) if total_error_count_row else 0 - ) + total_error_count = int(total_error_count_row[0]) if total_error_count_row else 0 return { "errors": [ @@ -1207,7 +1204,11 @@ class UsageAnalyticsStore: total_successful = int(row["total_successful"]) total_revenue_msats = float(row["total_revenue_msats"]) total_tokens = int(row["total_tokens"]) - if total_successful <= 0 and total_revenue_msats <= 0 and total_tokens <= 0: + if ( + total_successful <= 0 + and total_revenue_msats <= 0 + and total_tokens <= 0 + ): continue bucket_ts = str(row["bucket_ts"]) diff --git a/routstr/nostr/analytics.py b/routstr/nostr/analytics.py index 8b6da590..e568b5e0 100644 --- a/routstr/nostr/analytics.py +++ b/routstr/nostr/analytics.py @@ -215,9 +215,7 @@ def _build_window_payload( summary = dashboard.get("summary", {}) model_usage_mix = dashboard.get("model_usage_mix", {}) - summary_payload = _build_summary_payload( - summary if isinstance(summary, dict) else {} - ) + summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {}) usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {} top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload) @@ -340,9 +338,7 @@ async def publish_usage_analytics() -> None: nsec = (settings.nsec or "").strip() if not nsec: if not warned_missing_nsec: - logger.info( - "NSEC is not configured; skipping analytics sharing to Nostr" - ) + logger.info("NSEC is not configured; skipping analytics sharing to Nostr") warned_missing_nsec = True await asyncio.sleep(DISABLED_POLL_SECONDS) continue diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index e7cee8ca..37ac15d3 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -224,7 +224,9 @@ async def calculate_cost( "Token counts %s in the upstream response but cannot be " "priced; the request will appear in dashboards with the " "raw counts and a fixed max-cost charge.", - "are present" if (input_tokens > 0 or output_tokens > 0) else "are zero", + "are present" + if (input_tokens > 0 or output_tokens > 0) + else "are zero", extra={ "base_cost_msats": max_cost, "model": response_data.get("model", "unknown"), @@ -301,7 +303,9 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: # actually deducts from the balance. For non-BYOK providers (e.g. # OpenRouter) usage.cost already equals upstream_inference_cost, so we # fall through to the normal ``cost`` lookup below. - upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost")) + upstream_cost = _coerce_usd( + cost_details.get("upstream_inference_cost") + ) if upstream_cost > 0 and usage_data.get("is_byok"): byok_fee = _coerce_usd(usage_data.get("cost")) return upstream_cost + byok_fee @@ -332,7 +336,8 @@ def _get_pricing_rates( ``None`` means configured fixed pricing should be used by the caller. """ if settings.fixed_pricing and ( - settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens + settings.fixed_per_1k_input_tokens + or settings.fixed_per_1k_output_tokens ): return None @@ -388,8 +393,12 @@ def _get_pricing_rates( usd_per_sat = sats_usd_price() mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat - cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost")) - cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost")) + cache_read_usd = _coerce_usd( + pricing.get("cache_read_input_token_cost") + ) + cache_write_usd = _coerce_usd( + pricing.get("cache_creation_input_token_cost") + ) mscr_1k = ( cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat if cache_read_usd > 0 diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 11d5c01e..02c90055 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -110,7 +110,9 @@ def normalize_usage(usage_data: object) -> NormalizedUsage | None: if not isinstance(usage_data, dict): return None - output_tokens = _first_token_count(usage_data, "completion_tokens", "output_tokens") + output_tokens = _first_token_count( + usage_data, "completion_tokens", "output_tokens" + ) cache_read, cache_write = _extract_cache_tokens(usage_data) # ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index 985bcfd2..a693b763 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -94,7 +94,9 @@ class AzureUpstreamProvider(BaseUpstreamProvider): deployment_id = deployment_id.split("/")[-1] return f"openai/deployments/{deployment_id}/{clean_path}" - def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: + def get_request_base_url( + self, path: str, model_obj: "Model | None" = None + ) -> str: """Use endpoint root, stripping accidental /openai/v1 suffix if present.""" base_url = self.base_url.rstrip("/") marker = "/openai/v1" diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index a3d1505a..96955492 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -191,9 +191,7 @@ def _resolve_ehbp_target_url( otherwise the header is ignored so callers cannot redirect other providers or leak upstream API keys. """ - override_header = ( - profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER - ) + override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER if not override_header: return target_url enclave_url = _get_header_case_insensitive(headers, override_header) @@ -297,7 +295,9 @@ def _build_cost_info( return result -def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None: +def _inject_cost_response_headers( + headers: dict[str, str], cost_info: dict +) -> None: """Add per-request cost headers to an EHBP response. Since EHBP response bodies are opaque encrypted blobs, cost cannot be @@ -375,7 +375,9 @@ async def _compute_ehbp_actual_cost( resolved_upstream_model = ( actual_model_obj.forwarded_model_id or actual_model_obj.id ) - resolved_identity = _normalize_upstream_model_id(resolved_upstream_model) + resolved_identity = _normalize_upstream_model_id( + resolved_upstream_model + ) if resolved_identity != expected_identity: logger.info( "EHBP served model differs from requested, using actual " @@ -515,9 +517,7 @@ async def finalize_ehbp_actual_cost_payment( billing_key = await get_billing_key(key, session) key_hash = key.hashed_key billing_key_hash = billing_key.hashed_key - total_cost_msats = max( - 0, int(cost_info.get("total_msats", reserved_cost_for_model)) - ) + total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model))) now = int(time.time()) safe_reserved = case( @@ -560,9 +560,7 @@ async def finalize_ehbp_actual_cost_payment( ) child_result = await session.exec(child_stmt) # type: ignore[call-overload] - if result.rowcount == 0 or ( - child_result is not None and child_result.rowcount == 0 - ): + if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): await session.rollback() logger.error( "Failed to finalize EHBP usage-based payment", @@ -692,9 +690,7 @@ async def finalize_ehbp_max_cost_payment( else: child_result = None - if result.rowcount == 0 or ( - child_result is not None and child_result.rowcount == 0 - ): + if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): await session.rollback() logger.error( "Failed to finalize EHBP max-cost payment", @@ -1038,9 +1034,7 @@ async def forward_ehbp_x_cashu_request( target_url = _resolve_ehbp_target_url( target.url, path, headers, provider_type, profile ) - upstream_headers = _prepare_ehbp_upstream_headers( - headers, target.headers, profile - ) + upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile) request_body = await request.body() # Merge query params into the target URL @@ -1088,7 +1082,9 @@ async def forward_ehbp_x_cashu_request( usage_source = ( "header" if usage_header_name - and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers) + and any( + k.lower() == usage_header_name.lower() for k, _ in resp.headers + ) else ("trailer" if usage_header else "none") ) diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index d58199c2..54de41a3 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -94,7 +94,9 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): """ return self.base_url.rstrip("/").removesuffix("/openai") + "/openai" - def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: + def get_request_base_url( + self, path: str, model_obj: "Model | None" = None + ) -> str: """Route every proxied request to the OpenAI-compat surface. Required because the stored ``base_url`` typically points at the diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index a7e41c9b..11440b91 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -371,7 +371,9 @@ async def dispatch_gemini_messages( aggregates). """ if not request_body: - raise UpstreamError("Missing request body for /v1/messages", status_code=400) + raise UpstreamError( + "Missing request body for /v1/messages", status_code=400 + ) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index a0c9475e..17103c35 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -20,9 +20,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row( - cls, provider_row: "UpstreamProviderRow" - ) -> "GroqUpstreamProvider": + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/upstream/litellm_routing.py b/routstr/upstream/litellm_routing.py index b7790394..0b2a92a3 100644 --- a/routstr/upstream/litellm_routing.py +++ b/routstr/upstream/litellm_routing.py @@ -91,7 +91,9 @@ OLLAMA_HOST_HINTS: tuple[str, ...] = ( ) -def detect_litellm_prefix(base_url: str | None, default: str = DEFAULT_PREFIX) -> str: +def detect_litellm_prefix( + base_url: str | None, default: str = DEFAULT_PREFIX +) -> str: """Return the litellm provider prefix (`"/"`) for `base_url`. Falls back to `default` when the host doesn't match any known provider. diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index efcf591b..3d689922 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -108,7 +108,9 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: return events, buffer -def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]: +def events_from_chunk( + chunk: object, sse_buffer: bytes +) -> tuple[list[dict], bytes]: """Normalize a stream chunk into one or more event dicts. ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE @@ -199,7 +201,9 @@ async def aggregate_anthropic_events_to_message( raw_json = partial_json.pop(idx, None) if raw_json is not None and idx < len(blocks): try: - blocks[idx]["input"] = json.loads(raw_json) if raw_json else {} + blocks[idx]["input"] = ( + json.loads(raw_json) if raw_json else {} + ) except json.JSONDecodeError: blocks[idx]["input"] = raw_json elif etype == "message_delta": @@ -441,7 +445,9 @@ async def dispatch_anthropic_messages( on bad input or upstream failure. """ if not request_body: - raise UpstreamError("Missing request body for /v1/messages", status_code=400) + raise UpstreamError( + "Missing request body for /v1/messages", status_code=400 + ) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index beaafa6a..f3bb9971 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -1,19 +1,15 @@ """Model-path discovery service. -Exposes every upstream provider path a Routstr model is reachable through. -This is discovery/visibility data only — routing still selects the cheapest or -best provider separately. +Exposes every selectable upstream route a Routstr model is reachable through. +This PR remains discovery-only: request-side routing will consume the opaque +selectors in a follow-up. -A *path* is the provider string that may appear in Routstr chat completion -responses. The strings emitted here are produced by the provider's own -``discovery_path_for_subprovider`` / ``discovery_base_paths`` hooks, which -mirror ``_apply_provider_field`` so discovery and response stamping cannot -drift: +A path is a standard percent-encoded query string containing the normalized +upstream URL and, for an exact OpenRouter endpoint, its machine-readable tag. +Display names never participate in identity:: -- Direct upstream -> ```` e.g. ``anthropic`` -- Generic/custom OpenRouter-compatible upstream -> ``generic:`` -- Native OpenRouter routing to a sub-provider -> ``openrouter:`` -- Native OpenRouter with no usable sub-provider -> ``unknown`` + url=https%3A%2F%2Fapi.anthropic.com%2Fv1 + url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider=google-vertex%2Fus """ from __future__ import annotations @@ -21,10 +17,12 @@ from __future__ import annotations import asyncio import random import time -from typing import TYPE_CHECKING, Callable +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable +from urllib.parse import urlencode import httpx -from sqlalchemy import insert, or_ +from sqlalchemy import insert from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select @@ -51,6 +49,46 @@ _PERSIST_CHUNK_SIZE = 500 ModelKey = tuple[str, int] +@dataclass(frozen=True) +class EndpointIdentity: + """Exact OpenRouter endpoint identity returned by ``/endpoints``.""" + + tag: str + provider_name: str | None + + +@dataclass(frozen=True) +class DiscoveredPath: + """One model route ready for persistence and API serialization.""" + + model_id: str + path: str + upstream_url: str + provider_tag: str | None = None + provider_name: str | None = None + + +@dataclass(frozen=True) +class ProviderPathSnapshot: + """Refresh result plus model IDs whose prior rows must survive degradation.""" + + paths: tuple[DiscoveredPath, ...] + preserve_model_ids: frozenset[str] = frozenset() + + +def normalize_upstream_url(base_url: str) -> str: + """Normalize route identity without changing URL semantics.""" + return base_url.rstrip("/") + + +def encode_model_path(base_url: str, provider_tag: str | None = None) -> str: + """Encode a stable opaque selector for future request-side routing.""" + components = [("url", normalize_upstream_url(base_url))] + if provider_tag: + components.append(("provider", provider_tag)) + return urlencode(components) + + def _make_http_client() -> httpx.AsyncClient: """Client factory, separated so tests can substitute a mock transport.""" return httpx.AsyncClient() @@ -69,9 +107,15 @@ def is_openrouter_base_url(base_url: str | None) -> bool: def exposed_model_id(model: object) -> str: - """Client-visible ``/v1/models`` id for a cached model.""" + """Return exactly the ID advertised by ``/v1/models``. + + A forwarded ID is already a public routable alias and must remain intact, + including any slash. Without one, ``/v1/models`` exposes the base ID. + """ forwarded = getattr(model, "forwarded_model_id", None) - return forwarded or getattr(model, "id") + if forwarded: + return forwarded + return public_model_id(getattr(model, "id")) def public_model_id(model_id: str) -> str: @@ -116,7 +160,7 @@ class _RefreshCycleState: """ def __init__(self) -> None: - self.endpoint_cache: dict[tuple[str, str], list[str] | None] = {} + self.endpoint_cache: dict[tuple[str, str], list[EndpointIdentity] | None] = {} self.rate_limited = False @@ -127,8 +171,8 @@ async def _fetch_openrouter_endpoint_subproviders( author_slug: str, semaphore: asyncio.Semaphore, cycle: _RefreshCycleState, -) -> list[str] | None: - """Return sub-provider names for one model, or ``None`` when unknown. +) -> list[EndpointIdentity] | None: + """Return exact endpoint identities for one model, or ``None`` when unknown. ``None`` (not ``[]``) signals a degraded fetch — network failure, rate limit, non-200, or an unparseable payload — so callers can distinguish @@ -143,7 +187,7 @@ async def _fetch_openrouter_endpoint_subproviders( url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} - result: list[str] | None + result: list[EndpointIdentity] | None async with semaphore: try: resp = await client.get( @@ -177,14 +221,24 @@ async def _fetch_openrouter_endpoint_subproviders( endpoints = resp.json().get("data", {}).get("endpoints", []) if not isinstance(endpoints, list): endpoints = [] - names: list[str] = [] + identities: dict[str, EndpointIdentity] = {} for endpoint in endpoints: - provider_name = ( - endpoint.get("provider_name") if isinstance(endpoint, dict) else None + if not isinstance(endpoint, dict): + continue + tag = endpoint.get("tag") + if not isinstance(tag, str) or not tag.strip(): + continue + provider_name = endpoint.get("provider_name") + identities.setdefault( + tag, + EndpointIdentity( + tag=tag, + provider_name=provider_name + if isinstance(provider_name, str) and provider_name + else None, + ), ) - if provider_name: - names.append(provider_name) - result = list(dict.fromkeys(names)) + result = list(identities.values()) except Exception as e: # noqa: BLE001 logger.warning( "OpenRouter endpoint discovery bad payload", @@ -287,46 +341,41 @@ async def _collect_provider_paths( overrides_by_key: dict[ModelKey, ModelRow] | None = None, disabled_model_keys: set[ModelKey] | None = None, cycle: _RefreshCycleState | None = None, -) -> list[tuple[str, str]] | None: - """Collect ``(model_id, path)`` pairs for one provider instance. +) -> ProviderPathSnapshot: + """Collect selectable routes while marking model-level degraded fetches. - Emits the provider's ``discovery_base_paths`` for normal upstreams. For - OpenRouter-compatible providers, additionally emits one path per OpenRouter - sub-provider endpoint via ``discovery_path_for_subprovider`` so the strings - match response stamping exactly. - - Returns ``None`` when the provider's path set could not be determined this - cycle (every endpoint fetch degraded); callers must then keep previously - persisted rows instead of wiping them. + A failed OpenRouter lookup preserves only that model's prior rows. Other + models in the same provider still refresh, so a partial outage cannot erase + valid discovery data or freeze the entire provider snapshot. """ cycle = cycle or _RefreshCycleState() models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) - base_paths = upstream.discovery_base_paths() + upstream_url = normalize_upstream_url(upstream.base_url) + + def _base_path(model: object) -> DiscoveredPath: + return DiscoveredPath( + model_id=exposed_model_id(model), + path=encode_model_path(upstream_url), + upstream_url=upstream_url, + ) if not is_openrouter_base_url(upstream.base_url): - return [ - (exposed_model_id(model), path) for model in models for path in base_paths - ] + return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models)) if not (upstream.provider_type or "").strip(): - return [] + return ProviderPathSnapshot(paths=()) - any_fetch_succeeded = False - any_fetch_attempted = False semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) async with _make_http_client() as client: - async def _for_model(model: object) -> list[tuple[str, str]]: - nonlocal any_fetch_succeeded, any_fetch_attempted + async def _for_model( + model: object, + ) -> tuple[list[DiscoveredPath], str | None]: model_id = exposed_model_id(model) - # Base paths always apply: responses whose upstream payload lacks a - # provider field are stamped with them (see _apply_provider_field). - pairs = [(model_id, path) for path in base_paths] author_slug = openrouter_author_slug(model) if not author_slug: - return pairs - any_fetch_attempted = True - sub_providers = await _fetch_openrouter_endpoint_subproviders( + return [_base_path(model)], None + endpoints = await _fetch_openrouter_endpoint_subproviders( client, upstream.base_url, upstream.api_key, @@ -334,67 +383,78 @@ async def _collect_provider_paths( semaphore, cycle, ) - if sub_providers is None: - return [] - any_fetch_succeeded = True - paths = [ - upstream.discovery_path_for_subprovider(name) for name in sub_providers - ] - pairs.extend((model_id, path) for path in paths if path) - return list(dict.fromkeys(pairs)) + if endpoints is None: + return [], model_id + paths = [_base_path(model)] + paths.extend( + DiscoveredPath( + model_id=model_id, + path=encode_model_path(upstream_url, endpoint.tag), + upstream_url=upstream_url, + provider_tag=endpoint.tag, + provider_name=endpoint.provider_name, + ) + for endpoint in endpoints + ) + return paths, None results = await asyncio.gather( - *(_for_model(m) for m in models), return_exceptions=True + *(_for_model(model) for model in models), return_exceptions=True ) - if any_fetch_attempted and not any_fetch_succeeded: - # Every endpoint lookup degraded (offline, throttled, bad payloads): - # the true path set is unknown, not empty. - return None - - pairs: list[tuple[str, str]] = [] - for result in results: + paths: list[DiscoveredPath] = [] + preserve_model_ids: set[str] = set() + for model, result in zip(models, results): if isinstance(result, BaseException): + model_id = exposed_model_id(model) + preserve_model_ids.add(model_id) logger.warning( "OpenRouter endpoint discovery task errored", extra={"provider": upstream.provider_type, "error": str(result)}, ) continue - pairs.extend(result) + model_paths, preserved_model_id = result + paths.extend(model_paths) + if preserved_model_id: + preserve_model_ids.add(preserved_model_id) - return pairs + return ProviderPathSnapshot( + paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids) + ) async def _persist_provider_paths( - upstream_provider_id: int, pairs: list[tuple[str, str]] + upstream_provider_id: int, snapshot: ProviderPathSnapshot ) -> None: - """Replace all rows for ``upstream_provider_id`` with ``pairs``. - - Replacement (not upsert) so stale paths disappear when provider config or - upstream availability changes. Rows are written with chunked bulk INSERTs - so the transaction holds SQLite's write lock briefly — billing writes share - this database file. - """ - unique_pairs = list(dict.fromkeys(pairs)) + """Replace refreshed rows while retaining model-level degraded snapshots.""" + unique_paths = list( + {(path.model_id, path.path): path for path in snapshot.paths}.values() + ) now = int(time.time()) async with create_session() as session: - await session.exec( # type: ignore[call-overload] - delete(ModelPathRow).where( - col(ModelPathRow.upstream_provider_id) == upstream_provider_id - ) + delete_stmt = delete(ModelPathRow).where( + col(ModelPathRow.upstream_provider_id) == upstream_provider_id ) - for start in range(0, len(unique_pairs), _PERSIST_CHUNK_SIZE): - chunk = unique_pairs[start : start + _PERSIST_CHUNK_SIZE] + if snapshot.preserve_model_ids: + delete_stmt = delete_stmt.where( + col(ModelPathRow.model_id).not_in(sorted(snapshot.preserve_model_ids)) + ) + await session.exec(delete_stmt) # type: ignore[call-overload] + for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE): + chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE] await session.execute( insert(ModelPathRow), [ { - "model_id": model_id, - "path": path, + "model_id": discovered.model_id, + "path": discovered.path, + "upstream_url": discovered.upstream_url, + "provider_tag": discovered.provider_tag, + "provider_name": discovered.provider_name, "upstream_provider_id": upstream_provider_id, "updated_at": now, } - for model_id, path in chunk + for discovered in chunk ], ) await session.commit() @@ -454,22 +514,22 @@ async def refresh_model_paths( if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: continue try: - pairs = await _collect_provider_paths( + snapshot = await _collect_provider_paths( upstream, overrides_by_key=overrides_by_key, disabled_model_keys=disabled_model_keys, cycle=cycle, ) - if pairs is None: + if snapshot.preserve_model_ids: logger.warning( - "Model paths unknown this cycle; keeping previous rows", + "Some model paths are unknown; keeping their previous rows", extra={ "provider": upstream.provider_type or upstream.base_url, "db_id": upstream.db_id, + "preserved_models": len(snapshot.preserve_model_ids), }, ) - continue - await _persist_provider_paths(upstream.db_id, pairs) + await _persist_provider_paths(upstream.db_id, snapshot) except Exception as e: # noqa: BLE001 - isolate per-provider failures logger.error( "Failed to refresh model paths for provider", @@ -482,6 +542,21 @@ async def refresh_model_paths( ) +async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None: + """Immediately synchronize discovery after an admin provider/model mutation.""" + from ..proxy import get_upstreams + + matching = [ + upstream + for upstream in get_upstreams() + if upstream.db_id == upstream_provider_id + ] + if matching: + await refresh_model_paths(matching) + else: + await prune_model_paths_for_inactive_providers() + + def _refresh_interval_seconds() -> int: """Current interval, re-read every loop so runtime setting changes apply.""" from ..core.settings import settings @@ -535,8 +610,19 @@ async def refresh_model_paths_periodically( break +def _serialize_path(row: ModelPathRow) -> dict[str, Any]: + provider = None + if row.provider_tag or row.provider_name: + provider = {"name": row.provider_name, "slug": row.provider_tag} + return { + "path": row.path, + "upstream_url": row.upstream_url, + "provider": provider, + } + + async def get_all_model_paths() -> dict: - """All models with their paths, shaped for ``GET /v1/models/paths``.""" + """All models with their exact selectable routes.""" async with create_session() as session: rows = ( await session.exec( @@ -548,49 +634,37 @@ async def get_all_model_paths() -> dict: ) ).all() - grouped: dict[str, list[dict]] = {} + grouped: dict[str, list[dict[str, Any]]] = {} seen_paths: dict[str, set[str]] = {} updated_at = 0 for row in rows: updated_at = max(updated_at, row.updated_at) - model_id = public_model_id(row.model_id) - if row.path in seen_paths.setdefault(model_id, set()): + if row.path in seen_paths.setdefault(row.model_id, set()): continue - seen_paths[model_id].add(row.path) - grouped.setdefault(model_id, []).append({"path": row.path}) - # Deterministic output: models sorted by public id, paths sorted within. - data: list[dict] = [] - for grouped_model_id in sorted(grouped): - model_paths = sorted(grouped[grouped_model_id], key=lambda p: str(p["path"])) - data.append({"id": grouped_model_id, "paths": model_paths}) + seen_paths[row.model_id].add(row.path) + grouped.setdefault(row.model_id, []).append(_serialize_path(row)) + data = [ + { + "id": grouped_model_id, + "paths": sorted( + grouped[grouped_model_id], key=lambda item: str(item["path"]) + ), + } + for grouped_model_id in sorted(grouped) + ] return {"data": data, "updated_at": updated_at or None} async def get_paths_for_model(model_id: str) -> dict: - """Paths for a single model, shaped for ``GET /v1/models/paths/model``. - - Match by the public, unqualified model id, mirroring the model cache alias - behavior. Both ``deepseek-v4-pro`` and ``deepseek/deepseek-v4-pro`` resolve - every row whose stored id has the same base model id. The candidate set is - narrowed in SQL (exact id or ``%/`` suffix) so the route does not - materialize the whole table per request. - """ - # The request may be a full stored id ("z-ai/glm-5v-turbo") or an - # already-stripped public id ("fireworks/models/glm-5"); accept both. - accepted_ids = {model_id, public_model_id(model_id)} + """Return paths only for the exact model ID advertised by ``/v1/models``.""" async with create_session() as session: - conditions = [] - for candidate in accepted_ids: - conditions.append(col(ModelPathRow.model_id) == candidate) - conditions.append(col(ModelPathRow.model_id).endswith(f"/{candidate}")) rows = ( await session.exec( select(ModelPathRow) - .where(or_(*conditions)) + .where(col(ModelPathRow.model_id) == model_id) .order_by( col(ModelPathRow.path), col(ModelPathRow.upstream_provider_id), - col(ModelPathRow.model_id), ) ) ).all() @@ -599,15 +673,9 @@ async def get_paths_for_model(model_id: str) -> dict: paths: list[dict] = [] updated_at = 0 for row in rows: - # The SQL suffix match is a prefilter; enforce the exact public-id rule. - if ( - row.model_id not in accepted_ids - and public_model_id(row.model_id) not in accepted_ids - ): - continue updated_at = max(updated_at, row.updated_at) if row.path in seen: continue seen.add(row.path) - paths.append({"path": row.path}) + paths.append(_serialize_path(row)) return {"data": paths, "updated_at": updated_at or None} diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index c4873ea0..9fed0154 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -66,7 +66,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: + def get_request_base_url( + self, path: str, model_obj: Model | None = None + ) -> str: """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" return f"{self.base_url.rstrip('/')}/v1" @@ -183,9 +185,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): except Exception: self._models_cache = models_with_fees - self._models_by_id = { - m.forwarded_model_id or m.id: m for m in self._models_cache - } + self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache} logger.info( f"Refreshed models cache for {self.base_url}", extra={"model_count": len(models)}, diff --git a/routstr/upstream/rate_limit.py b/routstr/upstream/rate_limit.py index ac1eff78..dca1ba5b 100644 --- a/routstr/upstream/rate_limit.py +++ b/routstr/upstream/rate_limit.py @@ -119,9 +119,7 @@ def classify_rate_limit( retry_match = _RETRY_RE.search(redacted) if retry_match is not None: value = float(retry_match.group(1)) - retry_after = ( - value / 1000.0 if retry_match.group(2).lower() == "ms" else value - ) + retry_after = value / 1000.0 if retry_match.group(2).lower() == "ms" else value limit_name_match = _LIMIT_NAME_RE.search(redacted) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index 8e4379a0..c2ea5b1d 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -84,7 +84,9 @@ def extract_error_message(response: Response) -> str: return "" -def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: +def strip_unsupported_param( + body: dict, error_message: str +) -> tuple[dict, str] | None: """Drop a top-level param the upstream named as unsupported/deprecated. Returns ``(new_body, param)`` (a new dict, original untouched) when the diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index de1aa3bd..0371946a 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -50,7 +50,8 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): def normalize_request_path( self, path: str, model_obj: "Model | None" = None ) -> str: - """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.""" + """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr. + """ return path.lstrip("/") @classmethod diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 12e3dd93..58caaba0 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -21,9 +21,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row( - cls, provider_row: "UpstreamProviderRow" - ) -> "XAIUpstreamProvider": + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/wallet.py b/routstr/wallet.py index cef6902d..dd92d913 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -243,9 +243,7 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { - f"{k.mint_url}/{k.unit.name}": sum( - p.amount for p in wallet.proofs if p.id == k.id - ) + f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id) for k in wallet.keysets.values() } # Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet @@ -600,16 +598,11 @@ async def swap_to_primary_mint( # advance the counter so the next request derives fresh secrets. logger.warning( "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", - extra={ - "mint_quote_id": mint_quote.quote, - "minted_amount": minted_amount, - }, + extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, ) try: for keyset_id in primary_wallet.keysets: - await primary_wallet.restore_tokens_for_keyset( - keyset_id, to=1, batch=25 - ) + await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) await primary_wallet.load_proofs(reload=True) post_recovery_balance = primary_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance From 16fc548b48d982211f4f4a8afde63a6a36f6da9c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 20:12:55 +0200 Subject: [PATCH 09/14] fix: make model path identity selectable --- docs/api/endpoints.md | 56 +-- .../4e0c3d195a49_add_model_paths_table.py | 4 + routstr/core/admin.py | 28 ++ routstr/core/db.py | 14 +- routstr/payment/models.py | 10 +- routstr/upstream/model_paths.py | 107 +++--- tests/unit/test_model_paths.py | 347 +++++++++++++----- 7 files changed, 401 insertions(+), 165 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 44e6207e..2d202659 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -329,8 +329,8 @@ GET /v1/models ### List Model Paths -Get the upstream provider paths each advertised model can be reached through. -This is discovery data only; routing still chooses the provider per request. +Get the selectable upstream routes for each advertised model. This endpoint is +discovery-only; request-side selection will be added separately. ```http GET /v1/models/paths @@ -342,45 +342,45 @@ GET /v1/models/paths { "data": [ { - "id": "claude-sonnet-4", + "id": "anthropic/claude-sonnet-4", "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"} + { + "path": "provider=12", + "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, + "endpoint": null + }, + { + "path": "provider=42&endpoint=google-vertex%2Fus", + "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, + "endpoint": {"tag": "google-vertex/us", "name": "Google"} + } ] } - ] + ], + "updated_at": 1753500000 } ``` +`path` is an opaque, percent-encoded selector. Clients must store and return it +unchanged rather than parsing or reconstructing it. The configured provider's +stable node-local ID defines the upstream route; no upstream URL is exposed. +OpenRouter routes additionally use the exact machine-readable endpoint `tag`. +Provider slugs/types and endpoint names are display data and never participate +in identity. When request-side selection is implemented, an endpoint tag must +not silently fall back to another backend. + ### List Paths for One Model -Use a query parameter so model IDs containing `/` are handled safely. Lookup is -by the public, unqualified model ID: `glm-5v-turbo` resolves -`z-ai/glm-5v-turbo`, and `deepseek-v4-pro` and `deepseek/deepseek-v4-pro` -return the same merged path set. +Use the exact model ID advertised by `/v1/models`. The query parameter safely +supports IDs containing `/`. ```http GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 ``` -**Response:** - -```json -{ - "data": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"} - ] -} -``` - -Model IDs in responses are base model IDs: the leading provider prefix such as -`z-ai/` or `openai/` is stripped (the same rule routing uses, so the ID can be -sent back to `/v1/chat/completions` verbatim). Path values match the provider -string stamped on chat-completion responses, such as `anthropic`, -`generic:Anthropic`, `openrouter:Anthropic`, or `unknown` (native OpenRouter -with no usable sub-provider). Responses also carry an `updated_at` Unix -timestamp of the last successful refresh (`null` when no refresh has run). +The response uses the same path objects and `updated_at` field as the collection +endpoint. An unknown model returns `404 Model not found`. A known model whose +paths have not been discovered yet returns `200` with an empty `data` array. ## Wallet Management diff --git a/migrations/versions/4e0c3d195a49_add_model_paths_table.py b/migrations/versions/4e0c3d195a49_add_model_paths_table.py index 5641e488..688e2bc5 100644 --- a/migrations/versions/4e0c3d195a49_add_model_paths_table.py +++ b/migrations/versions/4e0c3d195a49_add_model_paths_table.py @@ -22,6 +22,10 @@ def upgrade() -> None: sa.Column("id", sa.Integer(), nullable=False), sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("provider_slug", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("provider_type", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("endpoint_tag", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("endpoint_name", sqlmodel.sql.sqltypes.AutoString(), nullable=True), sa.Column("upstream_provider_id", sa.Integer(), nullable=False), sa.Column("updated_at", sa.Integer(), nullable=False, server_default="0"), sa.ForeignKeyConstraint( diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 66a1d288..f581b704 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -51,6 +51,23 @@ ADMIN_SESSION_DURATION = 3600 MAX_USAGE_ANALYTICS_HOURS = 365 * 24 +async def _refresh_provider_model_paths(upstream_provider_id: int) -> None: + """Best-effort immediate discovery sync after an admin mutation.""" + from ..upstream.model_paths import refresh_model_paths_for_provider + + try: + await refresh_model_paths_for_provider(upstream_provider_id) + except Exception as exc: # noqa: BLE001 - committed admin writes must survive + logger.warning( + "Failed to refresh model paths after admin mutation", + extra={ + "upstream_provider_id": upstream_provider_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) + + async def require_admin_api(request: Request) -> None: auth_header = request.headers.get("Authorization") if not auth_header or not auth_header.startswith("Bearer "): @@ -579,6 +596,7 @@ async def upsert_provider_model( await session.refresh(row) await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return _row_to_model( row, apply_provider_fee=True, provider_fee=provider.provider_fee ).dict() # type: ignore @@ -633,6 +651,7 @@ async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, ob await session.delete(row) await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return {"ok": True, "deleted_id": model_id} @@ -652,6 +671,7 @@ async def delete_all_provider_models(provider_id: str) -> dict[str, object]: await session.delete(row) # type: ignore await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return {"ok": True, "deleted": len(rows)} @@ -705,6 +725,9 @@ async def batch_override_provider_models( json.dumps(model_data.alias_ids) if model_data.alias_ids else None ) existing_row.enabled = model_data.enabled + existing_row.forwarded_model_id = ( + model_data.forwarded_model_id or model_data.id + ) session.add(existing_row) else: # Create new @@ -735,6 +758,7 @@ async def batch_override_provider_models( ), upstream_provider_id=provider_pk, enabled=model_data.enabled, + forwarded_model_id=model_data.forwarded_model_id or model_data.id, ) session.add(row) @@ -743,6 +767,7 @@ async def batch_override_provider_models( await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return { "ok": True, "count": overridden_count, @@ -943,6 +968,7 @@ async def create_upstream_provider( await reinitialize_upstreams() await refresh_model_maps() + await _refresh_provider_model_paths(_provider_pk(provider)) return _serialize_provider(provider) @@ -968,6 +994,7 @@ async def update_upstream_provider( await reinitialize_upstreams() await refresh_model_maps() + await _refresh_provider_model_paths(_provider_pk(provider)) return _serialize_provider(provider) @@ -1003,6 +1030,7 @@ async def update_upstream_provider_by_slug( await reinitialize_upstreams() await refresh_model_maps() + await _refresh_provider_model_paths(_provider_pk(provider)) return _serialize_provider(provider) diff --git a/routstr/core/db.py b/routstr/core/db.py index 8b3586ec..151d8627 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -338,8 +338,18 @@ class ModelPathRow(SQLModel, table=True): # type: ignore description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( - description="Provider path stamped on chat completion responses, e.g. " - "'anthropic' or 'openrouter:Anthropic'" + description="Opaque selector containing provider ID and optional endpoint tag" + ) + provider_slug: str = Field( + description="Public slug of the configured upstream provider" + ) + provider_type: str = Field(description="Configured upstream provider type") + endpoint_tag: str | None = Field( + default=None, + description="Exact OpenRouter endpoint tag used for request-side selection", + ) + endpoint_name: str | None = Field( + default=None, description="Human-readable endpoint display name" ) upstream_provider_id: int = Field( index=True, diff --git a/routstr/payment/models.py b/routstr/payment/models.py index c433ddfa..5c634ced 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -613,9 +613,17 @@ async def model_paths_for_model(model_id: str) -> dict: model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL encoding and there is no dynamic-route ambiguity. """ + from ..proxy import get_unique_models from ..upstream.model_paths import get_paths_for_model - return await get_paths_for_model(model_id) + result = await get_paths_for_model(model_id) + if not result["data"]: + advertised_ids = { + model.forwarded_model_id or model.id for model in get_unique_models() + } + if model_id not in advertised_ids: + raise HTTPException(status_code=404, detail="Model not found") + return result @models_router.get("/v1/models") diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index f3bb9971..45ce6887 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -4,12 +4,13 @@ Exposes every selectable upstream route a Routstr model is reachable through. This PR remains discovery-only: request-side routing will consume the opaque selectors in a follow-up. -A path is a standard percent-encoded query string containing the normalized -upstream URL and, for an exact OpenRouter endpoint, its machine-readable tag. -Display names never participate in identity:: +A path is a standard percent-encoded query string containing the configured +provider's stable node-local ID and, for an exact OpenRouter endpoint, its +machine-readable tag. Upstream URLs and display names never participate in or +leak through public identity:: - url=https%3A%2F%2Fapi.anthropic.com%2Fv1 - url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider=google-vertex%2Fus + provider=42 + provider=42&endpoint=google-vertex%2Fus """ from __future__ import annotations @@ -57,15 +58,24 @@ class EndpointIdentity: provider_name: str | None +@dataclass(frozen=True) +class ConfiguredProviderIdentity: + """Public-safe identity of one configured upstream provider.""" + + id: int + slug: str + provider_type: str + + @dataclass(frozen=True) class DiscoveredPath: """One model route ready for persistence and API serialization.""" model_id: str path: str - upstream_url: str - provider_tag: str | None = None - provider_name: str | None = None + provider: ConfiguredProviderIdentity + endpoint_tag: str | None = None + endpoint_name: str | None = None @dataclass(frozen=True) @@ -76,16 +86,11 @@ class ProviderPathSnapshot: preserve_model_ids: frozenset[str] = frozenset() -def normalize_upstream_url(base_url: str) -> str: - """Normalize route identity without changing URL semantics.""" - return base_url.rstrip("/") - - -def encode_model_path(base_url: str, provider_tag: str | None = None) -> str: - """Encode a stable opaque selector for future request-side routing.""" - components = [("url", normalize_upstream_url(base_url))] - if provider_tag: - components.append(("provider", provider_tag)) +def encode_model_path(provider_id: int, endpoint_tag: str | None = None) -> str: + """Encode a stable opaque selector without exposing upstream URLs.""" + components: list[tuple[str, str | int]] = [("provider", provider_id)] + if endpoint_tag: + components.append(("endpoint", endpoint_tag)) return urlencode(components) @@ -218,9 +223,11 @@ async def _fetch_openrouter_endpoint_subproviders( return None try: - endpoints = resp.json().get("data", {}).get("endpoints", []) + 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(endpoints, list): - endpoints = [] + raise ValueError("endpoints must be a list") identities: dict[str, EndpointIdentity] = {} for endpoint in endpoints: if not isinstance(endpoint, dict): @@ -238,6 +245,8 @@ async def _fetch_openrouter_endpoint_subproviders( else None, ), ) + if endpoints and not identities: + raise ValueError("endpoints contain no usable tags") result = list(identities.values()) except Exception as e: # noqa: BLE001 logger.warning( @@ -251,7 +260,9 @@ async def _fetch_openrouter_endpoint_subproviders( async def _load_model_visibility() -> tuple[ - dict[ModelKey, ModelRow], set[ModelKey], set[int] + dict[ModelKey, ModelRow], + set[ModelKey], + dict[int, ConfiguredProviderIdentity], ]: """Load the same DB model visibility inputs used by routing. @@ -271,12 +282,16 @@ async def _load_model_visibility() -> tuple[ overrides_by_key: dict[ModelKey, ModelRow] = {} disabled_model_keys: set[ModelKey] = set() - enabled_provider_ids: set[int] = set() + provider_identities: dict[int, ConfiguredProviderIdentity] = {} for provider in provider_rows: if not provider.enabled or provider.id is None: continue - enabled_provider_ids.add(provider.id) + provider_identities[provider.id] = ConfiguredProviderIdentity( + id=provider.id, + slug=provider.slug or f"provider-{provider.id}", + provider_type=provider.provider_type, + ) for model in provider.models: key = (model.id.lower(), provider.id) if model.enabled: @@ -284,7 +299,7 @@ async def _load_model_visibility() -> tuple[ else: disabled_model_keys.add(key) - return overrides_by_key, disabled_model_keys, enabled_provider_ids + return overrides_by_key, disabled_model_keys, provider_identities def _apply_model_visibility( @@ -338,6 +353,7 @@ def _apply_model_visibility( async def _collect_provider_paths( upstream: BaseUpstreamProvider, + provider_identity: ConfiguredProviderIdentity, overrides_by_key: dict[ModelKey, ModelRow] | None = None, disabled_model_keys: set[ModelKey] | None = None, cycle: _RefreshCycleState | None = None, @@ -350,13 +366,12 @@ async def _collect_provider_paths( """ cycle = cycle or _RefreshCycleState() models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) - upstream_url = normalize_upstream_url(upstream.base_url) def _base_path(model: object) -> DiscoveredPath: return DiscoveredPath( model_id=exposed_model_id(model), - path=encode_model_path(upstream_url), - upstream_url=upstream_url, + path=encode_model_path(provider_identity.id), + provider=provider_identity, ) if not is_openrouter_base_url(upstream.base_url): @@ -389,10 +404,10 @@ async def _collect_provider_paths( paths.extend( DiscoveredPath( model_id=model_id, - path=encode_model_path(upstream_url, endpoint.tag), - upstream_url=upstream_url, - provider_tag=endpoint.tag, - provider_name=endpoint.provider_name, + path=encode_model_path(provider_identity.id, endpoint.tag), + provider=provider_identity, + endpoint_tag=endpoint.tag, + endpoint_name=endpoint.provider_name, ) for endpoint in endpoints ) @@ -448,9 +463,10 @@ async def _persist_provider_paths( { "model_id": discovered.model_id, "path": discovered.path, - "upstream_url": discovered.upstream_url, - "provider_tag": discovered.provider_tag, - "provider_name": discovered.provider_name, + "provider_slug": discovered.provider.slug, + "provider_type": discovered.provider.provider_type, + "endpoint_tag": discovered.endpoint_tag, + "endpoint_name": discovered.endpoint_name, "upstream_provider_id": upstream_provider_id, "updated_at": now, } @@ -505,17 +521,18 @@ async def refresh_model_paths( ( overrides_by_key, disabled_model_keys, - enabled_provider_ids, + provider_identities, ) = await _load_model_visibility() await prune_model_paths_for_inactive_providers() cycle = _RefreshCycleState() for upstream in upstreams: - if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: + if upstream.db_id is None or upstream.db_id not in provider_identities: continue try: snapshot = await _collect_provider_paths( upstream, + provider_identity=provider_identities[upstream.db_id], overrides_by_key=overrides_by_key, disabled_model_keys=disabled_model_keys, cycle=cycle, @@ -611,13 +628,17 @@ async def refresh_model_paths_periodically( def _serialize_path(row: ModelPathRow) -> dict[str, Any]: - provider = None - if row.provider_tag or row.provider_name: - provider = {"name": row.provider_name, "slug": row.provider_tag} + endpoint = None + if row.endpoint_tag or row.endpoint_name: + endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name} return { "path": row.path, - "upstream_url": row.upstream_url, - "provider": provider, + "provider": { + "id": row.upstream_provider_id, + "slug": row.provider_slug, + "type": row.provider_type, + }, + "endpoint": endpoint, } @@ -646,9 +667,7 @@ async def get_all_model_paths() -> dict: data = [ { "id": grouped_model_id, - "paths": sorted( - grouped[grouped_model_id], key=lambda item: str(item["path"]) - ), + "paths": grouped[grouped_model_id], } for grouped_model_id in sorted(grouped) ] diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 7d3393fe..0538ea55 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -142,11 +142,18 @@ def _mock_transport( return counter -def _endpoints_response(*provider_names: str) -> httpx.Response: - return httpx.Response( - 200, - json={"data": {"endpoints": [{"provider_name": n} for n in provider_names]}}, - ) +def _endpoints_response( + *providers: str | tuple[str, str], +) -> httpx.Response: + endpoints = [] + for provider in providers: + if isinstance(provider, tuple): + provider_name, tag = provider + else: + provider_name = provider + tag = provider.lower().replace(" ", "-") + endpoints.append({"provider_name": provider_name, "tag": tag}) + return httpx.Response(200, json={"data": {"endpoints": endpoints}}) _SEEDED_PROVIDER_IDS = (1, 2, 4, 5, 7) @@ -206,6 +213,29 @@ def _ids_of(payload: dict) -> set[str]: return {entry["id"] for entry in payload["data"]} +def _path_entry( + provider_id: int, + *, + provider_slug: str | None = None, + provider_type: str | None = None, + endpoint_tag: str | None = None, + endpoint_name: str | None = None, +) -> dict[str, Any]: + endpoint = None + if endpoint_tag or endpoint_name: + endpoint = {"tag": endpoint_tag, "name": endpoint_name} + return { + "path": mp.encode_model_path(provider_id, endpoint_tag), + "provider": { + "id": provider_id, + "slug": provider_slug or f"p{provider_id}", + "type": provider_type + or ("anthropic" if provider_id == 1 else "openrouter"), + }, + "endpoint": endpoint, + } + + # --------------------------------------------------------------------------- # # Predicates / pure helpers # --------------------------------------------------------------------------- # @@ -223,6 +253,13 @@ def test_native_anthropic_not_openrouter() -> None: assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False +def test_encode_model_path_uses_provider_id_without_exposing_url() -> None: + assert mp.encode_model_path(42) == "provider=42" + assert mp.encode_model_path(42, "google-vertex/us-east5") == ( + "provider=42&endpoint=google-vertex%2Fus-east5" + ) + + def test_exposed_model_id_prefers_forwarded() -> None: assert ( mp.exposed_model_id(_model("claude-x", forwarded_model_id="fwd-claude")) @@ -300,9 +337,7 @@ async def test_direct_provider_single_path_uses_provider_type( ) await mp.refresh_model_paths([provider]) payload = await mp.get_all_model_paths() - assert payload["data"] == [ - {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} - ] + assert payload["data"] == [{"id": "claude-opus-4.6", "paths": [_path_entry(1)]}] assert payload["updated_at"] is not None @@ -320,6 +355,22 @@ async def test_direct_path_stores_exposed_model_id( assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} +@pytest.mark.asyncio +async def test_forwarded_model_id_with_slash_remains_exact_and_routable( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("local-alias", forwarded_model_id="anthropic/claude-opus-4.6")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _ids_of(await mp.get_all_model_paths()) == {"anthropic/claude-opus-4.6"} + assert (await mp.get_paths_for_model("anthropic/claude-opus-4.6"))["data"] + + @pytest.mark.asyncio async def test_disabled_cached_models_excluded( patched_session: AsyncEngine, @@ -361,7 +412,7 @@ async def test_disabling_model_on_one_provider_keeps_other_provider( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {"anthropic"} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} @pytest.mark.asyncio @@ -395,8 +446,8 @@ async def test_override_alias_not_applied_across_providers( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {"anthropic"} - assert _paths_of(payload, "private-alias") == {"generic"} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} + assert _paths_of(payload, "private-alias") == {mp.encode_model_path(2)} @pytest.mark.asyncio @@ -464,7 +515,7 @@ 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": "anthropic"}]} + {"id": "public-alias", "paths": [_path_entry(1)]} ] @@ -486,7 +537,7 @@ 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": "generic"}]} + {"id": "public-deployment", "paths": [_path_entry(1)]} ] @@ -638,24 +689,34 @@ async def test_openrouter_provider_adds_endpoint_paths( ) _mock_transport( monkeypatch, - lambda request: _endpoints_response("Anthropic", "Amazon Bedrock"), + lambda request: _endpoints_response( + ("Google", "google-vertex/eu"), + ("Google", "google-vertex/us"), + ), ) await mp.refresh_model_paths([provider]) - paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert "openrouter:Anthropic" in paths - assert "openrouter:Amazon Bedrock" in paths - assert "openrouter" not in paths + payload = await mp.get_paths_for_model("claude-opus-4.6") + assert {item["path"] for item in payload["data"]} == { + mp.encode_model_path(2), + mp.encode_model_path(2, "google-vertex/eu"), + mp.encode_model_path(2, "google-vertex/us"), + } + assert { + item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] + } == {"google-vertex/eu", "google-vertex/us"} + assert { + item["endpoint"]["name"] for item in payload["data"] if item["endpoint"] + } == {"Google"} + assert {item["provider"]["id"] for item in payload["data"]} == {2} @pytest.mark.asyncio -async def test_openrouter_self_echoing_subprovider_maps_to_unknown( +async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch ) -> None: - """Responses stamped for a sub-provider echoing "OpenRouter" say - ``unknown``; discovery must advertise the same string, never - ``openrouter:OpenRouter``.""" + """Machine-readable endpoint tags, not display names, define identity.""" provider = _FakeOpenRouterProvider( models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], db_id=2, @@ -665,17 +726,14 @@ async def test_openrouter_self_echoing_subprovider_maps_to_unknown( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert "openrouter:OpenRouter" not in paths - assert "unknown" in paths + assert paths == {mp.encode_model_path(2), mp.encode_model_path(2, "openrouter")} @pytest.mark.asyncio async def test_generic_provider_with_openrouter_base_url_discovers( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch ) -> None: - """A generic provider pointed at OpenRouter exposes both the bare - ``generic`` path (stamped when the upstream omits its provider field) and - the ``generic:`` endpoint paths.""" + """Configured provider identity is independent from its endpoint URL.""" provider = _FakeProvider( provider_type="generic", base_url="https://openrouter.ai/api/v1", @@ -687,7 +745,35 @@ async def test_generic_provider_with_openrouter_base_url_discovers( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert paths == {"generic", "generic:Anthropic"} + assert paths == {mp.encode_model_path(1), mp.encode_model_path(1, "anthropic")} + + +@pytest.mark.asyncio +async def test_openrouter_partial_failure_keeps_failed_models_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + provider = _FakeOpenRouterProvider( + models=[ + _model("good", canonical_slug="author/good"), + _model("degraded", canonical_slug="author/degraded"), + ], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "degraded") + assert before + + def _partial_failure(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/author/degraded/endpoints"): + return httpx.Response(503) + return _endpoints_response("Google") + + _mock_transport(monkeypatch, _partial_failure) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "degraded") == before + assert _paths_of(await mp.get_all_model_paths(), "good") != before @pytest.mark.asyncio @@ -703,7 +789,7 @@ async def test_openrouter_failure_keeps_previous_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert "openrouter:Anthropic" in before + assert mp.encode_model_path(2, "anthropic") in before def _network_down(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("network down", request=request) @@ -726,10 +812,8 @@ async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) - assert _paths_of(await mp.get_all_model_paths(), "m0") == { - "unknown", - "openrouter:Anthropic", - } + expected = {mp.encode_model_path(2), mp.encode_model_path(2, "anthropic")} + assert _paths_of(await mp.get_all_model_paths(), "m0") == expected counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) await mp.refresh_model_paths([provider]) @@ -739,27 +823,28 @@ async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( assert counter["requests"] <= mp._OPENROUTER_CONCURRENCY, ( "429 must abort the remaining fan-out" ) - assert _paths_of(await mp.get_all_model_paths(), "m0") == { - "unknown", - "openrouter:Anthropic", - } + assert _paths_of(await mp.get_all_model_paths(), "m0") == expected @pytest.mark.asyncio -async def test_openrouter_bad_payload_shapes_do_not_raise( +async def test_openrouter_bad_payload_shapes_preserve_previous_rows( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch ) -> None: - """``endpoints: null`` and non-list endpoint payloads are swallowed as - documented, not raised into the generic task-errored bucket.""" + """Malformed successful responses are degraded snapshots, not empty sets.""" + provider = _FakeOpenRouterProvider( + models=[_model("m", canonical_slug="a/m")], db_id=2 + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "m") + for payload in ( {"data": {"endpoints": None}}, {"data": {"endpoints": "none"}}, + {"data": {"endpoints": [{"provider_name": "Anthropic"}]}}, {"data": None}, {}, ): - provider = _FakeOpenRouterProvider( - models=[_model("m", canonical_slug="a/m")], db_id=2 - ) def _handler( request: httpx.Request, p: dict[str, Any] | None = payload @@ -767,8 +852,8 @@ async def test_openrouter_bad_payload_shapes_do_not_raise( return httpx.Response(200, json=p) _mock_transport(monkeypatch, _handler) - # Must not raise. await mp.refresh_model_paths([provider]) + assert _paths_of(await mp.get_all_model_paths(), "m") == before @pytest.mark.asyncio @@ -795,8 +880,12 @@ async def test_openrouter_shared_base_url_fetched_once( assert counter["requests"] == 1 paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert "openrouter:Anthropic" in paths - assert "generic:Anthropic" in paths + assert paths == { + mp.encode_model_path(2), + mp.encode_model_path(2, "anthropic"), + mp.encode_model_path(4), + mp.encode_model_path(4, "anthropic"), + } @pytest.mark.asyncio @@ -857,13 +946,16 @@ async def test_same_model_two_providers_two_paths( assert len(payload["data"]) == 1 entry = payload["data"][0] assert entry["id"] == "claude-opus-4.6" - assert {p["path"] for p in entry["paths"]} == {"anthropic", "generic"} + assert {p["path"] for p in entry["paths"]} == { + mp.encode_model_path(1), + mp.encode_model_path(2), + } assert "canonical_id" not in entry assert all("canonical_id" not in p for p in entry["paths"]) @pytest.mark.asyncio -async def test_get_all_model_paths_deduplicates_visible_paths( +async def test_get_all_model_paths_keeps_distinct_configured_providers( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -881,7 +973,10 @@ async def test_get_all_model_paths_deduplicates_visible_paths( await mp.refresh_model_paths([p1, p2]) assert (await mp.get_all_model_paths())["data"] == [ - {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + { + "id": "claude-opus-4.6", + "paths": [_path_entry(1), _path_entry(2)], + } ] @@ -906,14 +1001,13 @@ async def test_get_all_model_paths_is_deterministic( @pytest.mark.asyncio -async def test_get_paths_for_model_returns_only_paths( +async def test_get_paths_for_model_returns_route_identity( patched_session: AsyncEngine, ) -> None: await _seed_two_provider_shared_model(patched_session) payload = await mp.get_paths_for_model("claude-opus-4.6") - assert {p["path"] for p in payload["data"]} == {"anthropic", "generic"} - assert all(set(p.keys()) == {"path"} for p in payload["data"]) + assert payload["data"] == [_path_entry(1), _path_entry(2)] assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] @@ -929,13 +1023,11 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( ) await mp.refresh_model_paths([provider]) - assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ - {"path": "generic"} - ] + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [_path_entry(4)] @pytest.mark.asyncio -async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( +async def test_get_paths_for_model_requires_exact_advertised_id( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -955,8 +1047,8 @@ async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] - assert {p["path"] for p in short_paths} == {"generic", "anthropic"} - assert prefixed_paths == short_paths + assert short_paths == [_path_entry(4), _path_entry(7)] + assert prefixed_paths == [] @pytest.mark.asyncio @@ -975,18 +1067,40 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ - {"path": "generic"} + _path_entry(1) ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" - ] == [{"path": "generic"}] + ] == [] # --------------------------------------------------------------------------- # -# Periodic refresh loop +# Immediate and periodic refresh # --------------------------------------------------------------------------- # +@pytest.mark.asyncio +async def test_refresh_model_paths_for_provider_selects_mutated_provider( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + target = SimpleNamespace(db_id=2) + other = SimpleNamespace(db_id=1) + seen: list[list[Any]] = [] + + monkeypatch.setattr(proxy, "get_upstreams", lambda: [other, target]) + + async def _fake_refresh(upstreams: list[Any]) -> None: + seen.append(upstreams) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + await mp.refresh_model_paths_for_provider(2) + + assert seen == [[target]] + + @pytest.mark.asyncio async def test_refresh_loop_rereads_interval_and_picks_up_providers( monkeypatch: pytest.MonkeyPatch, @@ -1093,37 +1207,88 @@ def _make_model_paths_app() -> FastAPI: def test_model_paths_endpoint_returns_all_paths( monkeypatch: pytest.MonkeyPatch, ) -> None: + expected = { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + _path_entry(1), + _path_entry( + 2, + endpoint_tag="google-vertex/us", + endpoint_name="Google", + ), + ], + } + ], + "updated_at": 1753500000, + } + async def _fake_get_all_model_paths() -> dict[str, Any]: - return { - "data": [ - { - "id": "claude-opus-4.6", - "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"}, - ], - } - ], - "updated_at": 1753500000, - } + return expected monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) response = TestClient(_make_model_paths_app()).get("/v1/models/paths") assert response.status_code == 200 - assert response.json() == { - "data": [ - { - "id": "claude-opus-4.6", - "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"}, - ], - } - ], - "updated_at": 1753500000, - } + assert response.json() == expected + + +def test_model_paths_for_model_returns_404_for_unknown_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: []) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "does-not-exist"} + ) + + assert response.status_code == 404 + assert response.json() == {"detail": "Model not found"} + + +def test_model_paths_for_known_model_can_return_empty_collection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: [_model("known")]) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "known"} + ) + + assert response.status_code == 200 + assert response.json() == {"data": [], "updated_at": None} + + +def test_model_paths_for_routing_only_alias_returns_404( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: [_model("advertised")]) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "routing-alias"} + ) + + assert response.status_code == 404 def test_model_paths_for_model_endpoint_accepts_slash_model_id( @@ -1131,9 +1296,14 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) -> None: calls: list[str] = [] + expected = { + "data": [_path_entry(2, endpoint_tag="anthropic", endpoint_name="Anthropic")], + "updated_at": None, + } + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: calls.append(model_id) - return {"data": [{"path": "generic:Anthropic"}], "updated_at": None} + return expected monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) @@ -1143,8 +1313,5 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) assert response.status_code == 200 - assert response.json() == { - "data": [{"path": "generic:Anthropic"}], - "updated_at": None, - } + assert response.json() == expected assert calls == ["anthropic/claude-opus-4.6"] From e2f89a26454debb3b44860508209457eaff50ab6 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 27 Jul 2026 23:17:13 +0200 Subject: [PATCH 10/14] fix: address follow-up model path review --- docs/api/endpoints.md | 6 +- routstr/core/admin.py | 20 +--- routstr/core/db.py | 2 +- routstr/upstream/base.py | 22 ---- routstr/upstream/model_paths.py | 114 +++++++++++++++---- routstr/upstream/openrouter.py | 17 --- tests/unit/test_fee_payout_migration.py | 10 +- tests/unit/test_model_paths.py | 144 ++++++++++++++++++------ 8 files changed, 213 insertions(+), 122 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 2d202659..ca54b8e0 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -345,12 +345,12 @@ GET /v1/models/paths "id": "anthropic/claude-sonnet-4", "paths": [ { - "path": "provider=12", + "path": "provider=anthropic-primary", "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, "endpoint": null }, { - "path": "provider=42&endpoint=google-vertex%2Fus", + "path": "provider=openrouter-main&endpoint=google-vertex%2Fus", "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, "endpoint": {"tag": "google-vertex/us", "name": "Google"} } @@ -363,7 +363,7 @@ GET /v1/models/paths `path` is an opaque, percent-encoded selector. Clients must store and return it unchanged rather than parsing or reconstructing it. The configured provider's -stable node-local ID defines the upstream route; no upstream URL is exposed. +public slug defines the upstream route; no upstream URL is exposed. OpenRouter routes additionally use the exact machine-readable endpoint `tag`. Provider slugs/types and endpoint names are display data and never participate in identity. When request-side selection is implemented, an endpoint tag must diff --git a/routstr/core/admin.py b/routstr/core/admin.py index f581b704..0402f1a7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -52,20 +52,10 @@ MAX_USAGE_ANALYTICS_HOURS = 365 * 24 async def _refresh_provider_model_paths(upstream_provider_id: int) -> None: - """Best-effort immediate discovery sync after an admin mutation.""" - from ..upstream.model_paths import refresh_model_paths_for_provider + """Queue discovery sync without blocking the committed admin mutation.""" + from ..upstream.model_paths import schedule_model_paths_refresh_for_provider - try: - await refresh_model_paths_for_provider(upstream_provider_id) - except Exception as exc: # noqa: BLE001 - committed admin writes must survive - logger.warning( - "Failed to refresh model paths after admin mutation", - extra={ - "upstream_provider_id": upstream_provider_id, - "error": str(exc), - "error_type": type(exc).__name__, - }, - ) + await schedule_model_paths_refresh_for_provider(upstream_provider_id) async def require_admin_api(request: Request) -> None: @@ -725,9 +715,6 @@ async def batch_override_provider_models( json.dumps(model_data.alias_ids) if model_data.alias_ids else None ) existing_row.enabled = model_data.enabled - existing_row.forwarded_model_id = ( - model_data.forwarded_model_id or model_data.id - ) session.add(existing_row) else: # Create new @@ -758,7 +745,6 @@ async def batch_override_provider_models( ), upstream_provider_id=provider_pk, enabled=model_data.enabled, - forwarded_model_id=model_data.forwarded_model_id or model_data.id, ) session.add(row) diff --git a/routstr/core/db.py b/routstr/core/db.py index 151d8627..3edfe140 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -338,7 +338,7 @@ class ModelPathRow(SQLModel, table=True): # type: ignore description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( - description="Opaque selector containing provider ID and optional endpoint tag" + description="Opaque selector containing provider slug and optional endpoint tag" ) provider_slug: str = Field( description="Public slug of the configured upstream provider" diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 5f9116c9..a8dba7e3 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -237,28 +237,6 @@ class BaseUpstreamProvider: except (TypeError, ValueError): pass - def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: - """Discovery path for a reported sub-provider name. - - Must produce exactly the value ``_apply_provider_field`` would stamp on - a response whose upstream payload reported ``sub_provider``, so the - model-path discovery API never advertises a path that cannot appear on - the wire. Subclasses that override ``_apply_provider_field`` must - override this to match. - """ - provider_type = (self.provider_type or "").strip() - if not provider_type: - return None - sub = (sub_provider or "").strip() - if not sub or sub == provider_type or sub.startswith(f"{provider_type}:"): - return sub or provider_type - return f"{provider_type}:{sub}" - - def discovery_base_paths(self) -> list[str]: - """Paths stamped when the upstream reports no sub-provider of its own.""" - provider_type = (self.provider_type or "").strip() - return [provider_type] if provider_type else [] - def _apply_provider_field(self, response_json: object) -> None: """Stamp the routstr ``provider`` field onto an upstream response payload. diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 45ce6887..40022aca 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -5,12 +5,12 @@ This PR remains discovery-only: request-side routing will consume the opaque selectors in a follow-up. A path is a standard percent-encoded query string containing the configured -provider's stable node-local ID and, for an exact OpenRouter endpoint, its +provider's public slug and, for an exact OpenRouter endpoint, its machine-readable tag. Upstream URLs and display names never participate in or leak through public identity:: - provider=42 - provider=42&endpoint=google-vertex%2Fus + provider=anthropic-primary + provider=openrouter-main&endpoint=google-vertex%2Fus """ from __future__ import annotations @@ -23,7 +23,7 @@ from typing import TYPE_CHECKING, Any, Callable from urllib.parse import urlencode import httpx -from sqlalchemy import insert +from sqlalchemy.dialects.sqlite import insert from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select @@ -44,6 +44,12 @@ _OPENROUTER_TIMEOUT_SECONDS = 10.0 # avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle. _PERSIST_CHUNK_SIZE = 500 +# Admin mutations enqueue provider IDs here instead of running OpenRouter's +# per-model endpoint fan-out inside the request. One worker serializes refreshes +# and coalesces repeated mutations for the same provider. +_scheduled_provider_refresh_ids: set[int] = set() +_scheduled_provider_refresh_task: asyncio.Task[None] | None = None + # Visibility key used across this module: routing carries the provider # dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)), # so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps. @@ -86,9 +92,9 @@ class ProviderPathSnapshot: preserve_model_ids: frozenset[str] = frozenset() -def encode_model_path(provider_id: int, endpoint_tag: str | None = None) -> str: +def encode_model_path(provider_slug: str, endpoint_tag: str | None = None) -> str: """Encode a stable opaque selector without exposing upstream URLs.""" - components: list[tuple[str, str | int]] = [("provider", provider_id)] + components = [("provider", provider_slug)] if endpoint_tag: components.append(("endpoint", endpoint_tag)) return urlencode(components) @@ -370,7 +376,7 @@ async def _collect_provider_paths( def _base_path(model: object) -> DiscoveredPath: return DiscoveredPath( model_id=exposed_model_id(model), - path=encode_model_path(provider_identity.id), + path=encode_model_path(provider_identity.slug), provider=provider_identity, ) @@ -404,7 +410,7 @@ async def _collect_provider_paths( paths.extend( DiscoveredPath( model_id=model_id, - path=encode_model_path(provider_identity.id, endpoint.tag), + path=encode_model_path(provider_identity.slug, endpoint.tag), provider=provider_identity, endpoint_tag=endpoint.tag, endpoint_name=endpoint.provider_name, @@ -457,21 +463,31 @@ async def _persist_provider_paths( await session.exec(delete_stmt) # type: ignore[call-overload] for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE): chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE] + values = [ + { + "model_id": discovered.model_id, + "path": discovered.path, + "provider_slug": discovered.provider.slug, + "provider_type": discovered.provider.provider_type, + "endpoint_tag": discovered.endpoint_tag, + "endpoint_name": discovered.endpoint_name, + "upstream_provider_id": upstream_provider_id, + "updated_at": now, + } + for discovered in chunk + ] + insert_stmt = insert(ModelPathRow).values(values) await session.execute( - insert(ModelPathRow), - [ - { - "model_id": discovered.model_id, - "path": discovered.path, - "provider_slug": discovered.provider.slug, - "provider_type": discovered.provider.provider_type, - "endpoint_tag": discovered.endpoint_tag, - "endpoint_name": discovered.endpoint_name, - "upstream_provider_id": upstream_provider_id, - "updated_at": now, - } - for discovered in chunk - ], + insert_stmt.on_conflict_do_update( + index_elements=["model_id", "path", "upstream_provider_id"], + set_={ + "provider_slug": insert_stmt.excluded.provider_slug, + "provider_type": insert_stmt.excluded.provider_type, + "endpoint_tag": insert_stmt.excluded.endpoint_tag, + "endpoint_name": insert_stmt.excluded.endpoint_name, + "updated_at": insert_stmt.excluded.updated_at, + }, + ) ) await session.commit() @@ -560,7 +576,10 @@ async def refresh_model_paths( async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None: - """Immediately synchronize discovery after an admin provider/model mutation.""" + """Synchronize one provider when model-path discovery is enabled.""" + if _refresh_interval_seconds() <= 0: + return + from ..proxy import get_upstreams matching = [ @@ -574,6 +593,55 @@ async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None: await prune_model_paths_for_inactive_providers() +async def _drain_scheduled_provider_refreshes() -> None: + """Serialize and coalesce model-path refreshes scheduled by admin writes.""" + global _scheduled_provider_refresh_task + + try: + # Let mutations in the same event-loop turn collapse into one refresh. + await asyncio.sleep(0) + while _scheduled_provider_refresh_ids: + if _refresh_interval_seconds() <= 0: + _scheduled_provider_refresh_ids.clear() + return + provider_id = min(_scheduled_provider_refresh_ids) + _scheduled_provider_refresh_ids.remove(provider_id) + try: + await refresh_model_paths_for_provider(provider_id) + except asyncio.CancelledError: + raise + except Exception as exc: # noqa: BLE001 - background best effort + logger.warning( + "Failed to refresh model paths after admin mutation", + extra={ + "upstream_provider_id": provider_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) + finally: + _scheduled_provider_refresh_task = None + + +async def schedule_model_paths_refresh_for_provider( + upstream_provider_id: int, +) -> None: + """Queue a non-blocking, coalesced refresh after an admin mutation.""" + global _scheduled_provider_refresh_task + + if _refresh_interval_seconds() <= 0: + return + _scheduled_provider_refresh_ids.add(upstream_provider_id) + if ( + _scheduled_provider_refresh_task is None + or _scheduled_provider_refresh_task.done() + ): + _scheduled_provider_refresh_task = asyncio.create_task( + _drain_scheduled_provider_refreshes(), + name="model-path-admin-refresh", + ) + + def _refresh_interval_seconds() -> int: """Current interval, re-read every loop so runtime setting changes apply.""" from ..core.settings import settings diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index fe92c4f9..1caeaa5c 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -18,23 +18,6 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): supports_anthropic_messages = True litellm_provider_prefix = "openrouter/" - def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: - """Mirror ``_apply_provider_field``: strip repeated prefixes, map a - missing or self-echoing sub-provider to the literal ``"unknown"``.""" - provider_type = (self.provider_type or "").strip() - sub = (sub_provider or "").strip() - prefix = f"{provider_type}:" - while sub.lower().startswith(prefix.lower()): - sub = sub[len(prefix) :].strip() - if not sub or sub.lower() == provider_type.lower(): - return "unknown" - return f"{provider_type}:{sub}" - - def discovery_base_paths(self) -> list[str]: - """Native OpenRouter never stamps a bare ``openrouter``; a response - with no sub-provider is stamped ``unknown``.""" - return ["unknown"] - def _apply_provider_field(self, response_json: object) -> None: """Stamp the ``provider`` field for OpenRouter responses. diff --git a/tests/unit/test_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index c1ce13c8..17be72ec 100644 --- a/tests/unit/test_fee_payout_migration.py +++ b/tests/unit/test_fee_payout_migration.py @@ -4,6 +4,9 @@ import subprocess import sys from pathlib import Path +from alembic.config import Config +from alembic.script import ScriptDirectory + def _run_alembic(root: Path, database_url: str, revision: str) -> None: env = os.environ.copy() @@ -37,9 +40,10 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None: "payout_in_progress_msats, payout_started_at FROM routstr_fees" ).fetchone() - # Head of the 7f2843d3f4e4 lineage: model-paths chains onto the fee-payout - # repair migration. - assert version == ("4e0c3d195a49",) + migration_config = Config(str(root / "alembic.ini")) + assert version == ( + ScriptDirectory.from_config(migration_config).get_current_head(), + ) assert { "id", "accumulated_msats", diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 0538ea55..70986dfd 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -225,7 +225,7 @@ def _path_entry( if endpoint_tag or endpoint_name: endpoint = {"tag": endpoint_tag, "name": endpoint_name} return { - "path": mp.encode_model_path(provider_id, endpoint_tag), + "path": mp.encode_model_path(provider_slug or f"p{provider_id}", endpoint_tag), "provider": { "id": provider_id, "slug": provider_slug or f"p{provider_id}", @@ -253,10 +253,10 @@ def test_native_anthropic_not_openrouter() -> None: assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False -def test_encode_model_path_uses_provider_id_without_exposing_url() -> None: - assert mp.encode_model_path(42) == "provider=42" - assert mp.encode_model_path(42, "google-vertex/us-east5") == ( - "provider=42&endpoint=google-vertex%2Fus-east5" +def test_encode_model_path_uses_provider_slug_without_exposing_url() -> None: + assert mp.encode_model_path("openrouter-main") == "provider=openrouter-main" + assert mp.encode_model_path("openrouter-main", "google-vertex/us-east5") == ( + "provider=openrouter-main&endpoint=google-vertex%2Fus-east5" ) @@ -305,21 +305,6 @@ def test_openrouter_author_slug_none_when_no_slash() -> None: assert mp.openrouter_author_slug(m) is None -def test_discovery_paths_mirror_response_stamping() -> None: - """The discovery hook and ``_apply_provider_field`` must agree.""" - generic = _FakeProvider(provider_type="generic", base_url="https://x", models=[]) - assert generic.discovery_path_for_subprovider("Anthropic") == "generic:Anthropic" - assert generic.discovery_base_paths() == ["generic"] - - native = _FakeOpenRouterProvider(models=[]) - assert native.discovery_path_for_subprovider("GMICloud") == "openrouter:GMICloud" - # Sub-provider echoing the router name is stamped "unknown" on responses. - assert native.discovery_path_for_subprovider("OpenRouter") == "unknown" - assert native.discovery_path_for_subprovider("openrouter:openrouter") == "unknown" - assert native.discovery_path_for_subprovider(None) == "unknown" - assert native.discovery_base_paths() == ["unknown"] - - # --------------------------------------------------------------------------- # # Refresh through the public entry point # --------------------------------------------------------------------------- # @@ -412,7 +397,7 @@ async def test_disabling_model_on_one_provider_keeps_other_provider( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path("p1")} @pytest.mark.asyncio @@ -446,8 +431,8 @@ async def test_override_alias_not_applied_across_providers( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} - assert _paths_of(payload, "private-alias") == {mp.encode_model_path(2)} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path("p1")} + assert _paths_of(payload, "private-alias") == {mp.encode_model_path("p2")} @pytest.mark.asyncio @@ -699,9 +684,9 @@ async def test_openrouter_provider_adds_endpoint_paths( payload = await mp.get_paths_for_model("claude-opus-4.6") assert {item["path"] for item in payload["data"]} == { - mp.encode_model_path(2), - mp.encode_model_path(2, "google-vertex/eu"), - mp.encode_model_path(2, "google-vertex/us"), + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "google-vertex/eu"), + mp.encode_model_path("p2", "google-vertex/us"), } assert { item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] @@ -726,7 +711,10 @@ async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert paths == {mp.encode_model_path(2), mp.encode_model_path(2, "openrouter")} + assert paths == { + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "openrouter"), + } @pytest.mark.asyncio @@ -745,7 +733,10 @@ async def test_generic_provider_with_openrouter_base_url_discovers( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert paths == {mp.encode_model_path(1), mp.encode_model_path(1, "anthropic")} + assert paths == { + mp.encode_model_path("p1"), + mp.encode_model_path("p1", "anthropic"), + } @pytest.mark.asyncio @@ -776,6 +767,37 @@ async def test_openrouter_partial_failure_keeps_failed_models_previous_rows( assert _paths_of(await mp.get_all_model_paths(), "good") != before +@pytest.mark.asyncio +async def test_partial_failure_upserts_collapsed_public_model_paths( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A degraded canonical sibling may preserve the same public path that a + successful sibling refreshes; persistence must merge instead of rolling back.""" + provider = _FakeOpenRouterProvider( + models=[ + _model("vendora/shared", canonical_slug="vendora/shared"), + _model("vendorb/shared", canonical_slug="vendorb/shared"), + ], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + + def _one_sibling_degrades(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/vendorb/shared/endpoints"): + return httpx.Response(503) + return _endpoints_response("Google") + + _mock_transport(monkeypatch, _one_sibling_degrades) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "shared") == { + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "anthropic"), + mp.encode_model_path("p2", "google"), + } + + @pytest.mark.asyncio async def test_openrouter_failure_keeps_previous_rows( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch @@ -789,7 +811,7 @@ async def test_openrouter_failure_keeps_previous_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert mp.encode_model_path(2, "anthropic") in before + assert mp.encode_model_path("p2", "anthropic") in before def _network_down(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("network down", request=request) @@ -812,7 +834,10 @@ async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) - expected = {mp.encode_model_path(2), mp.encode_model_path(2, "anthropic")} + expected = { + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "anthropic"), + } assert _paths_of(await mp.get_all_model_paths(), "m0") == expected counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) @@ -881,10 +906,10 @@ async def test_openrouter_shared_base_url_fetched_once( assert counter["requests"] == 1 paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path(2), - mp.encode_model_path(2, "anthropic"), - mp.encode_model_path(4), - mp.encode_model_path(4, "anthropic"), + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "anthropic"), + mp.encode_model_path("p4"), + mp.encode_model_path("p4", "anthropic"), } @@ -947,8 +972,8 @@ async def test_same_model_two_providers_two_paths( entry = payload["data"][0] assert entry["id"] == "claude-opus-4.6" assert {p["path"] for p in entry["paths"]} == { - mp.encode_model_path(1), - mp.encode_model_path(2), + mp.encode_model_path("p1"), + mp.encode_model_path("p2"), } assert "canonical_id" not in entry assert all("canonical_id" not in p for p in entry["paths"]) @@ -1101,6 +1126,53 @@ async def test_refresh_model_paths_for_provider_selects_mutated_provider( assert seen == [[target]] +@pytest.mark.asyncio +async def test_admin_refresh_is_disabled_by_model_paths_kill_switch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + calls: list[int] = [] + + async def _fake_refresh(provider_id: int) -> None: + calls.append(provider_id) + + monkeypatch.setattr(mp, "refresh_model_paths_for_provider", _fake_refresh) + await mp.schedule_model_paths_refresh_for_provider(2) + await asyncio.sleep(0) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_admin_refresh_is_backgrounded_and_coalesced( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + mp._scheduled_provider_refresh_ids.clear() + mp._scheduled_provider_refresh_task = None + calls: list[int] = [] + + async def _fake_refresh(provider_id: int) -> None: + calls.append(provider_id) + + monkeypatch.setattr(mp, "refresh_model_paths_for_provider", _fake_refresh) + + await mp.schedule_model_paths_refresh_for_provider(2) + await mp.schedule_model_paths_refresh_for_provider(2) + task = mp._scheduled_provider_refresh_task + assert task is not None + await task + + assert calls == [2] + + @pytest.mark.asyncio async def test_refresh_loop_rereads_interval_and_picks_up_providers( monkeypatch: pytest.MonkeyPatch, From bb2a05b67ce0bb21c5b055a56041be9ea53ea6e0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 29 Jul 2026 00:08:34 +0200 Subject: [PATCH 11/14] add url and specific model infos to paht --- docs/api/endpoints.md | 17 ++-- routstr/core/db.py | 5 +- routstr/upstream/model_paths.py | 65 ++++++++++--- tests/unit/test_model_paths.py | 167 +++++++++++++++++++++++++------- 4 files changed, 195 insertions(+), 59 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index ca54b8e0..23b019c7 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -345,12 +345,12 @@ GET /v1/models/paths "id": "anthropic/claude-sonnet-4", "paths": [ { - "path": "provider=anthropic-primary", + "path": "url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=anthropic%2Fclaude-sonnet-4", "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, "endpoint": null }, { - "path": "provider=openrouter-main&endpoint=google-vertex%2Fus", + "path": "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4&endpoint=google-vertex%2Fus", "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, "endpoint": {"tag": "google-vertex/us", "name": "Google"} } @@ -362,12 +362,13 @@ GET /v1/models/paths ``` `path` is an opaque, percent-encoded selector. Clients must store and return it -unchanged rather than parsing or reconstructing it. The configured provider's -public slug defines the upstream route; no upstream URL is exposed. -OpenRouter routes additionally use the exact machine-readable endpoint `tag`. -Provider slugs/types and endpoint names are display data and never participate -in identity. When request-side selection is implemented, an endpoint tag must -not silently fall back to another backend. +unchanged rather than parsing or reconstructing it. It identifies the exact +configured route with `url`, `provider-id`, and `model-id`. To avoid exposing +private network details, a configured private IP address or any URL with an +explicit port is advertised as `http://localhost`. OpenRouter routes additionally +preserve the exact machine-readable endpoint `tag`. Provider slugs/types and +endpoint names remain display data. When request-side selection is implemented, +an endpoint tag must not silently fall back to another backend. ### List Paths for One Model diff --git a/routstr/core/db.py b/routstr/core/db.py index 3edfe140..e4794d4a 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -338,7 +338,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( - description="Opaque selector containing provider slug and optional endpoint tag" + description=( + "Opaque selector containing upstream URL, provider ID, model ID, " + "and optional endpoint tag" + ) ) provider_slug: str = Field( description="Public slug of the configured upstream provider" diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 40022aca..020f0fbb 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -5,22 +5,22 @@ This PR remains discovery-only: request-side routing will consume the opaque selectors in a follow-up. A path is a standard percent-encoded query string containing the configured -provider's public slug and, for an exact OpenRouter endpoint, its -machine-readable tag. Upstream URLs and display names never participate in or -leak through public identity:: +upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter +endpoint, its machine-readable tag:: - provider=anthropic-primary - provider=openrouter-main&endpoint=google-vertex%2Fus + url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=claude-sonnet-4 + url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=claude-sonnet-4&endpoint=google-vertex%2Fus """ from __future__ import annotations import asyncio +import ipaddress import random import time from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Callable -from urllib.parse import urlencode +from urllib.parse import urlencode, urlsplit import httpx from sqlalchemy.dialects.sqlite import insert @@ -66,11 +66,12 @@ class EndpointIdentity: @dataclass(frozen=True) class ConfiguredProviderIdentity: - """Public-safe identity of one configured upstream provider.""" + """Public identity of one configured upstream provider.""" id: int slug: str provider_type: str + base_url: str @dataclass(frozen=True) @@ -92,9 +93,38 @@ class ProviderPathSnapshot: preserve_model_ids: frozenset[str] = frozenset() -def encode_model_path(provider_slug: str, endpoint_tag: str | None = None) -> str: - """Encode a stable opaque selector without exposing upstream URLs.""" - components = [("provider", provider_slug)] +def public_provider_url(base_url: str) -> str: + """Mask private IP addresses and URLs with explicit ports.""" + parsed = urlsplit(base_url) + try: + if parsed.port is not None: + return "http://localhost" + except ValueError: + # An invalid explicit port must not accidentally leak through. + return "http://localhost" + + hostname = parsed.hostname + if hostname is None: + return base_url + try: + address = ipaddress.ip_address(hostname) + except ValueError: + return base_url + return "http://localhost" if address.is_private else base_url + + +def encode_model_path( + base_url: str, + provider_id: int, + model_id: str, + endpoint_tag: str | None = None, +) -> str: + """Encode the complete upstream route selector advertised to clients.""" + components: list[tuple[str, str | int]] = [ + ("url", base_url), + ("provider-id", provider_id), + ("model-id", model_id), + ] if endpoint_tag: components.append(("endpoint", endpoint_tag)) return urlencode(components) @@ -297,6 +327,7 @@ async def _load_model_visibility() -> tuple[ id=provider.id, slug=provider.slug or f"provider-{provider.id}", provider_type=provider.provider_type, + base_url=public_provider_url(provider.base_url), ) for model in provider.models: key = (model.id.lower(), provider.id) @@ -374,9 +405,12 @@ async def _collect_provider_paths( models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) def _base_path(model: object) -> DiscoveredPath: + model_id = exposed_model_id(model) return DiscoveredPath( - model_id=exposed_model_id(model), - path=encode_model_path(provider_identity.slug), + model_id=model_id, + path=encode_model_path( + provider_identity.base_url, provider_identity.id, model_id + ), provider=provider_identity, ) @@ -410,7 +444,12 @@ async def _collect_provider_paths( paths.extend( DiscoveredPath( model_id=model_id, - path=encode_model_path(provider_identity.slug, endpoint.tag), + path=encode_model_path( + provider_identity.base_url, + provider_identity.id, + model_id, + endpoint.tag, + ), provider=provider_identity, endpoint_tag=endpoint.tag, endpoint_name=endpoint.provider_name, diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 70986dfd..1bbe25e0 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -213,8 +213,19 @@ def _ids_of(payload: dict) -> set[str]: return {entry["id"] for entry in payload["data"]} +def _expected_path( + provider_id: int, + model_id: str, + endpoint_tag: str | None = None, +) -> str: + return mp.encode_model_path( + f"https://provider-{provider_id}", provider_id, model_id, endpoint_tag + ) + + def _path_entry( provider_id: int, + model_id: str, *, provider_slug: str | None = None, provider_type: str | None = None, @@ -225,7 +236,7 @@ def _path_entry( if endpoint_tag or endpoint_name: endpoint = {"tag": endpoint_tag, "name": endpoint_name} return { - "path": mp.encode_model_path(provider_slug or f"p{provider_id}", endpoint_tag), + "path": _expected_path(provider_id, model_id, endpoint_tag), "provider": { "id": provider_id, "slug": provider_slug or f"p{provider_id}", @@ -253,10 +264,38 @@ def test_native_anthropic_not_openrouter() -> None: assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False -def test_encode_model_path_uses_provider_slug_without_exposing_url() -> None: - assert mp.encode_model_path("openrouter-main") == "provider=openrouter-main" - assert mp.encode_model_path("openrouter-main", "google-vertex/us-east5") == ( - "provider=openrouter-main&endpoint=google-vertex%2Fus-east5" +def test_public_provider_url_masks_private_addresses_and_explicit_ports() -> None: + assert mp.public_provider_url("http://192.168.1.10/v1") == "http://localhost" + assert mp.public_provider_url("http://10.0.0.5:11434/v1") == "http://localhost" + assert mp.public_provider_url("http://[fd00::1]/v1") == "http://localhost" + assert mp.public_provider_url("https://api.example.com:8443/v1") == ( + "http://localhost" + ) + + +def test_public_provider_url_preserves_public_urls_without_ports() -> None: + assert mp.public_provider_url("https://openrouter.ai/api/v1") == ( + "https://openrouter.ai/api/v1" + ) + assert mp.public_provider_url("http://localhost") == "http://localhost" + + +def test_encode_model_path_includes_complete_route_identity() -> None: + assert mp.encode_model_path( + "https://openrouter.ai/api/v1", 42, "anthropic/claude-sonnet-4" + ) == ( + "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1" + "&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4" + ) + assert mp.encode_model_path( + "https://openrouter.ai/api/v1", + 42, + "anthropic/claude-sonnet-4", + "google-vertex/us-east5", + ) == ( + "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1" + "&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4" + "&endpoint=google-vertex%2Fus-east5" ) @@ -322,10 +361,36 @@ async def test_direct_provider_single_path_uses_provider_type( ) await mp.refresh_model_paths([provider]) payload = await mp.get_all_model_paths() - assert payload["data"] == [{"id": "claude-opus-4.6", "paths": [_path_entry(1)]}] + assert payload["data"] == [ + {"id": "claude-opus-4.6", "paths": [_path_entry(1, "claude-opus-4.6")]} + ] assert payload["updated_at"] is not None +@pytest.mark.asyncio +async def test_direct_path_masks_private_configured_provider_url( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.base_url = "http://192.168.1.10:11434/v1" + session.add(provider_row) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="http://192.168.1.10:11434/v1", + models=[_model("local-model")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "local-model") == { + mp.encode_model_path("http://localhost", 1, "local-model") + } + + @pytest.mark.asyncio async def test_direct_path_stores_exposed_model_id( patched_session: AsyncEngine, @@ -397,7 +462,9 @@ async def test_disabling_model_on_one_provider_keeps_other_provider( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {mp.encode_model_path("p1")} + assert _paths_of(payload, "shared-model") == { + _expected_path(1, "shared-model") + } @pytest.mark.asyncio @@ -431,8 +498,12 @@ async def test_override_alias_not_applied_across_providers( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {mp.encode_model_path("p1")} - assert _paths_of(payload, "private-alias") == {mp.encode_model_path("p2")} + assert _paths_of(payload, "shared-model") == { + _expected_path(1, "shared-model") + } + assert _paths_of(payload, "private-alias") == { + _expected_path(2, "private-alias") + } @pytest.mark.asyncio @@ -500,7 +571,7 @@ 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)]} + {"id": "public-alias", "paths": [_path_entry(1, "public-alias")]} ] @@ -522,7 +593,10 @@ 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)]} + { + "id": "public-deployment", + "paths": [_path_entry(1, "public-deployment")], + } ] @@ -684,9 +758,9 @@ async def test_openrouter_provider_adds_endpoint_paths( payload = await mp.get_paths_for_model("claude-opus-4.6") assert {item["path"] for item in payload["data"]} == { - mp.encode_model_path("p2"), - mp.encode_model_path("p2", "google-vertex/eu"), - mp.encode_model_path("p2", "google-vertex/us"), + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "google-vertex/eu"), + _expected_path(2, "claude-opus-4.6", "google-vertex/us"), } assert { item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] @@ -712,8 +786,8 @@ async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path("p2"), - mp.encode_model_path("p2", "openrouter"), + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "openrouter"), } @@ -734,8 +808,8 @@ async def test_generic_provider_with_openrouter_base_url_discovers( paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path("p1"), - mp.encode_model_path("p1", "anthropic"), + _expected_path(1, "claude-opus-4.6"), + _expected_path(1, "claude-opus-4.6", "anthropic"), } @@ -792,9 +866,9 @@ async def test_partial_failure_upserts_collapsed_public_model_paths( await mp.refresh_model_paths([provider]) assert _paths_of(await mp.get_all_model_paths(), "shared") == { - mp.encode_model_path("p2"), - mp.encode_model_path("p2", "anthropic"), - mp.encode_model_path("p2", "google"), + _expected_path(2, "shared"), + _expected_path(2, "shared", "anthropic"), + _expected_path(2, "shared", "google"), } @@ -811,7 +885,7 @@ async def test_openrouter_failure_keeps_previous_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert mp.encode_model_path("p2", "anthropic") in before + assert _expected_path(2, "claude-opus-4.6", "anthropic") in before def _network_down(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("network down", request=request) @@ -835,8 +909,8 @@ async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) expected = { - mp.encode_model_path("p2"), - mp.encode_model_path("p2", "anthropic"), + _expected_path(2, "m0"), + _expected_path(2, "m0", "anthropic"), } assert _paths_of(await mp.get_all_model_paths(), "m0") == expected @@ -906,10 +980,10 @@ async def test_openrouter_shared_base_url_fetched_once( assert counter["requests"] == 1 paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path("p2"), - mp.encode_model_path("p2", "anthropic"), - mp.encode_model_path("p4"), - mp.encode_model_path("p4", "anthropic"), + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "anthropic"), + _expected_path(4, "claude-opus-4.6"), + _expected_path(4, "claude-opus-4.6", "anthropic"), } @@ -972,8 +1046,8 @@ async def test_same_model_two_providers_two_paths( entry = payload["data"][0] assert entry["id"] == "claude-opus-4.6" assert {p["path"] for p in entry["paths"]} == { - mp.encode_model_path("p1"), - mp.encode_model_path("p2"), + _expected_path(1, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6"), } assert "canonical_id" not in entry assert all("canonical_id" not in p for p in entry["paths"]) @@ -1000,7 +1074,10 @@ async def test_get_all_model_paths_keeps_distinct_configured_providers( assert (await mp.get_all_model_paths())["data"] == [ { "id": "claude-opus-4.6", - "paths": [_path_entry(1), _path_entry(2)], + "paths": [ + _path_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ], } ] @@ -1032,7 +1109,10 @@ async def test_get_paths_for_model_returns_route_identity( await _seed_two_provider_shared_model(patched_session) payload = await mp.get_paths_for_model("claude-opus-4.6") - assert payload["data"] == [_path_entry(1), _path_entry(2)] + assert payload["data"] == [ + _path_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ] assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] @@ -1048,7 +1128,9 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( ) await mp.refresh_model_paths([provider]) - assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [_path_entry(4)] + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ + _path_entry(4, "glm-5v-turbo") + ] @pytest.mark.asyncio @@ -1072,7 +1154,10 @@ async def test_get_paths_for_model_requires_exact_advertised_id( short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] - assert short_paths == [_path_entry(4), _path_entry(7)] + assert short_paths == [ + _path_entry(4, "deepseek-v4-pro"), + _path_entry(7, "deepseek-v4-pro"), + ] assert prefixed_paths == [] @@ -1092,7 +1177,7 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ - _path_entry(1) + _path_entry(1, "fireworks/models/glm-5") ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" @@ -1284,9 +1369,10 @@ def test_model_paths_endpoint_returns_all_paths( { "id": "claude-opus-4.6", "paths": [ - _path_entry(1), + _path_entry(1, "claude-opus-4.6"), _path_entry( 2, + "claude-opus-4.6", endpoint_tag="google-vertex/us", endpoint_name="Google", ), @@ -1369,7 +1455,14 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( calls: list[str] = [] expected = { - "data": [_path_entry(2, endpoint_tag="anthropic", endpoint_name="Anthropic")], + "data": [ + _path_entry( + 2, + "anthropic/claude-opus-4.6", + endpoint_tag="anthropic", + endpoint_name="Anthropic", + ) + ], "updated_at": None, } From 60566313dc415facadf7f338d503008c4909fa73 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 31 Jul 2026 01:17:31 +0200 Subject: [PATCH 12/14] fix: return persisted API-key refund token --- routstr/balance.py | 32 +++++++++++++++++ tests/unit/test_balance.py | 72 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 104 insertions(+) diff --git a/routstr/balance.py b/routstr/balance.py index 91b19ce5..cf44e612 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -259,6 +259,36 @@ async def _lookup_key_no_create( return None +async def _get_persisted_api_key_refund( + key: ApiKey, session: AsyncSession +) -> dict[str, str] | None: + result = await session.exec( + select(CashuTransaction) + .where( + CashuTransaction.api_key_hashed_key == key.hashed_key, + CashuTransaction.type == "out", + CashuTransaction.source == "apikey", + ) + .order_by(col(CashuTransaction.created_at).desc()) + ) + refund = result.first() + if refund is None: + return None + if refund.swept: + raise HTTPException(status_code=410, detail="Refund has been swept") + + refund.collected = True + session.add(refund) + await session.commit() + + persisted = {"token": refund.token} + if refund.unit == "sat": + persisted["sats"] = str(refund.amount) + else: + persisted["msats"] = str(refund.amount) + return persisted + + async def _restore_balance( session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str ) -> None: @@ -353,6 +383,8 @@ async def refund_wallet_endpoint( if key.total_balance <= 0: if cached := await _refund_cache_get(bearer_value): return cached + if persisted := await _get_persisted_api_key_refund(key, session): + return persisted if key.parent_key_hash: raise HTTPException( diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 609e2557..f3775e22 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -221,6 +221,78 @@ def _make_api_key( return key +@pytest.mark.asyncio +async def test_apikey_refund_returns_persisted_token_after_cache_loss() -> None: + key = _make_api_key(balance=0, refund_currency="sat") + refund_token = "cashuApersisted_refund_token" + refund_tx = _make_cashu_tx( + token=refund_token, + amount=5, + unit="sat", + type="out", + request_id=None, + ) + refund_tx.source = "apikey" + refund_tx.api_key_hashed_key = key.hashed_key + + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.exec = AsyncMock(return_value=_exec_result(refund_tx)) + session.add = MagicMock() + session.commit = AsyncMock() + + with ( + patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), + patch("routstr.balance.send_token", AsyncMock()) as mock_send_token, + ): + result = await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + assert result == {"token": refund_token, "sats": "5"} + assert refund_tx.collected is True + session.add.assert_called_once_with(refund_tx) + session.commit.assert_awaited_once() + mock_send_token.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_apikey_refund_rejects_persisted_token_after_sweep() -> None: + from fastapi import HTTPException + + key = _make_api_key(balance=0, refund_currency="sat") + refund_tx = _make_cashu_tx( + token="cashuAswept_apikey_refund", + amount=5, + unit="sat", + request_id=None, + swept=True, + ) + refund_tx.source = "apikey" + refund_tx.api_key_hashed_key = key.hashed_key + + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.exec = AsyncMock(return_value=_exec_result(refund_tx)) + session.add = MagicMock() + session.commit = AsyncMock() + + with patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + assert exc_info.value.status_code == 410 + assert exc_info.value.detail == "Refund has been swept" + session.add.assert_not_called() + session.commit.assert_not_awaited() + + @pytest.mark.asyncio async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> None: key = _make_api_key(balance=5000, refund_currency="sat") From 19236ecc9db82a05dcaec6a4f988888bf5e5b5b5 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 31 Jul 2026 02:14:50 +0200 Subject: [PATCH 13/14] fix(db): increase default connection pool capacity --- .env.example | 6 +++--- routstr/core/settings.py | 14 +++++++------- tests/unit/test_settings.py | 12 ++++++------ 3 files changed, 16 insertions(+), 16 deletions(-) diff --git a/.env.example b/.env.example index 9f0c0bbc..17ae6951 100644 --- a/.env.example +++ b/.env.example @@ -26,9 +26,9 @@ ROUTSTR_SECRET_KEY= # logged at startup. Keep total capacity across all workers below the database # connection limit. Pre-ping is automatic for networked backends; SQLite may # explicitly opt in if desired. -# DATABASE_POOL_SIZE=5 -# DATABASE_MAX_OVERFLOW=10 -# DATABASE_POOL_TIMEOUT=30 +# DATABASE_POOL_SIZE=10 +# DATABASE_MAX_OVERFLOW=20 +# DATABASE_POOL_TIMEOUT=15 # DATABASE_POOL_RECYCLE=1800 # DATABASE_POOL_PRE_PING=false # Warn when a checkout is held this many seconds. diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 0caffd16..4674733a 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -106,14 +106,14 @@ class Settings(BaseSettings): default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS" ) - # Database connection-pool controls (advanced). Capacity defaults match - # SQLAlchemy's established queue-pool behavior. Pre-ping is enabled by the - # engine factory for networked backends; SQLite can explicitly opt in. - # These fields are env-only below. - database_pool_size: int = Field(default=5, ge=1, env="DATABASE_POOL_SIZE") - database_max_overflow: int = Field(default=10, ge=0, env="DATABASE_MAX_OVERFLOW") + # Database connection-pool controls (advanced). Capacity defaults provide + # headroom for Routstr's concurrent request and background-payment workload. + # Pre-ping is enabled by the engine factory for networked backends; SQLite + # can explicitly opt in. These fields are env-only below. + database_pool_size: int = Field(default=10, ge=1, env="DATABASE_POOL_SIZE") + database_max_overflow: int = Field(default=20, ge=0, env="DATABASE_MAX_OVERFLOW") database_pool_timeout: float = Field( - default=30.0, gt=0, env="DATABASE_POOL_TIMEOUT" + default=15.0, gt=0, env="DATABASE_POOL_TIMEOUT" ) database_pool_recycle: int = Field(default=1800, ge=0, env="DATABASE_POOL_RECYCLE") database_pool_pre_ping: bool = Field(default=False, env="DATABASE_POOL_PRE_PING") diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index 98a9e24f..34fa7101 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -62,11 +62,11 @@ def test_payout_settings_have_sensible_defaults() -> None: assert s.payout_interval_seconds == 900 -def test_database_pool_defaults_match_sqlalchemy_capacity() -> None: +def test_database_pool_defaults_provide_concurrency_headroom() -> None: s = Settings() - assert s.database_pool_size == 5 - assert s.database_max_overflow == 10 - assert s.database_pool_timeout == 30.0 + assert s.database_pool_size == 10 + assert s.database_max_overflow == 20 + assert s.database_pool_timeout == 15.0 assert s.database_pool_recycle == 1800 assert s.database_pool_pre_ping is False assert s.database_pool_hold_warn_seconds == 10.0 @@ -139,7 +139,7 @@ async def test_update_does_not_apply_env_only_fields_to_live_settings( from the running pool. """ monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False) - monkeypatch.setattr(settings, "database_pool_size", 5) + monkeypatch.setattr(settings, "database_pool_size", 10) engine = create_async_engine("sqlite+aiosqlite:///:memory:") async with AsyncSession(engine, expire_on_commit=False) as session: @@ -151,7 +151,7 @@ async def test_update_does_not_apply_env_only_fields_to_live_settings( # A non-env-only field still updates normally... assert settings.name == "PoolTweaker" # ...but the env-only pool size stays at the boot value. - assert settings.database_pool_size == 5 + assert settings.database_pool_size == 10 # ...and it is never written to the settings blob. blob = await _read_settings_blob(session) assert "database_pool_size" not in blob From e903aa3a9f37210b36e6851654d6b72a94d9a926 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 2 Aug 2026 23:16:01 +0200 Subject: [PATCH 14/14] clean up --- routstr/upstream/model_paths.py | 34 ++++++++++++++++++++++----------- tests/unit/test_model_paths.py | 21 +++++++------------- 2 files changed, 30 insertions(+), 25 deletions(-) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 020f0fbb..95b6edf2 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -31,6 +31,8 @@ from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_sessio from ..core.logging import get_logger if TYPE_CHECKING: + from sqlmodel.ext.asyncio.session import AsyncSession + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -782,18 +784,28 @@ async def get_all_model_paths() -> dict: async def get_paths_for_model(model_id: str) -> dict: - """Return paths only for the exact model ID advertised by ``/v1/models``.""" - async with create_session() as session: - rows = ( - await session.exec( - select(ModelPathRow) - .where(col(ModelPathRow.model_id) == model_id) - .order_by( - col(ModelPathRow.path), - col(ModelPathRow.upstream_provider_id), + """Return paths for an advertised ID or its provider-prefixed alias.""" + + async def load_rows(session: AsyncSession, lookup_id: str) -> list[ModelPathRow]: + return list( + ( + await session.exec( + select(ModelPathRow) + .where(col(ModelPathRow.model_id) == lookup_id) + .order_by( + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) ) - ) - ).all() + ).all() + ) + + async with create_session() as session: + rows = await load_rows(session, model_id) + if not rows: + unprefixed_id = public_model_id(model_id) + if unprefixed_id != model_id: + rows = await load_rows(session, unprefixed_id) seen: set[str] = set() paths: list[dict] = [] diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 1bbe25e0..240a4199 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -462,9 +462,7 @@ async def test_disabling_model_on_one_provider_keeps_other_provider( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == { - _expected_path(1, "shared-model") - } + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} @pytest.mark.asyncio @@ -498,12 +496,8 @@ async def test_override_alias_not_applied_across_providers( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == { - _expected_path(1, "shared-model") - } - assert _paths_of(payload, "private-alias") == { - _expected_path(2, "private-alias") - } + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} + assert _paths_of(payload, "private-alias") == {_expected_path(2, "private-alias")} @pytest.mark.asyncio @@ -1134,7 +1128,7 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( @pytest.mark.asyncio -async def test_get_paths_for_model_requires_exact_advertised_id( +async def test_get_paths_for_model_accepts_provider_prefixed_alias( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -1158,15 +1152,14 @@ async def test_get_paths_for_model_requires_exact_advertised_id( _path_entry(4, "deepseek-v4-pro"), _path_entry(7, "deepseek-v4-pro"), ] - assert prefixed_paths == [] + assert prefixed_paths == short_paths @pytest.mark.asyncio async def test_get_paths_for_model_multi_segment_id_matches_models_listing( patched_session: AsyncEngine, ) -> None: - """For three-segment ids the discovery id must be the same base id the - rest of the system exposes (first-slash rule), not the last segment.""" + """Three-segment upstream IDs resolve to the same first-slash public ID.""" provider = _FakeProvider( provider_type="generic", base_url="https://x/v1", @@ -1181,7 +1174,7 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" - ] == [] + ] == [_path_entry(1, "fireworks/models/glm-5")] # --------------------------------------------------------------------------- #