From f96acbb99cc5dbc11e64f9f933d8520031aaf059 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:23:31 +0200 Subject: [PATCH] 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"]