diff --git a/.env.example b/.env.example index 3d12ee5c..6ec85d71 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/docs/api/endpoints.md b/docs/api/endpoints.md index 5a9ea68d..23b019c7 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -327,6 +327,62 @@ GET /v1/models } ``` +### List Model Paths + +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 +``` + +**Response:** + +```json +{ + "data": [ + { + "id": "anthropic/claude-sonnet-4", + "paths": [ + { + "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": "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"} + } + ] + } + ], + "updated_at": 1753500000 +} +``` + +`path` is an opaque, percent-encoded selector. Clients must store and return it +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 + +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 +``` + +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 ### 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 eee669c4..11cc2934 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -146,6 +146,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 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` | Mint HTTP 429 responses create a per-mint cooldown. Operations that already hold Routstr's wallet mutation lock fail fast during that cooldown instead of waiting @@ -184,3 +186,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/4f2a4f3f62e0_add_model_paths_table.py b/migrations/versions/4f2a4f3f62e0_add_model_paths_table.py new file mode 100644 index 00000000..59fee690 --- /dev/null +++ b/migrations/versions/4f2a4f3f62e0_add_model_paths_table.py @@ -0,0 +1,52 @@ +"""add model paths table + +Revision ID: 4f2a4f3f62e0 +Revises: bf76270b66c4 +Create Date: 2026-08-02 23:28:24.760061 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "4f2a4f3f62e0" +down_revision = "bf76270b66c4" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + 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("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), + 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_upstream_provider_id"), + "model_paths", + ["upstream_provider_id"], + unique=False, + ) + + +def downgrade() -> None: + op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths") + op.drop_table("model_paths") diff --git a/routstr/balance.py b/routstr/balance.py index b106aebd..121e2659 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -312,6 +312,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, @@ -414,6 +444,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/routstr/core/admin.py b/routstr/core/admin.py index 74a94668..f7883261 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -51,6 +51,13 @@ ADMIN_SESSION_DURATION = 3600 MAX_USAGE_ANALYTICS_HOURS = 365 * 24 +async def _refresh_provider_model_paths(upstream_provider_id: int) -> None: + """Queue discovery sync without blocking the committed admin mutation.""" + from ..upstream.model_paths import schedule_model_paths_refresh_for_provider + + await schedule_model_paths_refresh_for_provider(upstream_provider_id) + + 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 +586,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 +641,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 +661,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)} @@ -743,6 +753,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 +954,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 +980,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 +1016,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 33c9c104..4185791d 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -317,9 +317,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) ) @@ -373,6 +371,60 @@ 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) + # 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( + description="Client-visible /v1/models id (forwarded_model_id or id)" + ) + path: str = Field( + 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" + ) + 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, + foreign_key="upstream_providers.id", + 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 __tablename__ = "lightning_invoices" @@ -666,9 +718,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())) @@ -767,9 +817,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 ca1a5a93..979f5cdb 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 @@ -130,6 +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()) + # 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) + ) payout_task = asyncio.create_task(periodic_payout()) if global_settings.nsec: nip91_task = asyncio.create_task(announce_provider()) @@ -137,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()) @@ -176,6 +182,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: @@ -209,6 +217,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: @@ -245,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 @@ -321,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 4366fad5..6852eaa2 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -110,8 +110,14 @@ 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") + 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" @@ -120,14 +126,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/routstr/payment/models.py b/routstr/payment/models.py index 17e8ac64..5c634ced 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), @@ -595,6 +595,37 @@ 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 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 ..proxy import get_unique_models + from ..upstream.model_paths import get_paths_for_model + + 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") @models_router.get("/v1/models/", include_in_schema=False) @models_router.get("/models") diff --git a/routstr/proxy.py b/routstr/proxy.py index 9b7d4077..4542d008 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -190,6 +190,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/model_paths.py b/routstr/upstream/model_paths.py new file mode 100644 index 00000000..95b6edf2 --- /dev/null +++ b/routstr/upstream/model_paths.py @@ -0,0 +1,819 @@ +"""Model-path discovery service. + +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 configured +upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter +endpoint, its machine-readable tag:: + + 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, urlsplit + +import httpx +from sqlalchemy.dialects.sqlite import insert +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 + +if TYPE_CHECKING: + from sqlmodel.ext.asyncio.session import AsyncSession + + 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 + +# 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 + +# 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. +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 ConfiguredProviderIdentity: + """Public identity of one configured upstream provider.""" + + id: int + slug: str + provider_type: str + base_url: str + + +@dataclass(frozen=True) +class DiscoveredPath: + """One model route ready for persistence and API serialization.""" + + model_id: str + path: str + provider: ConfiguredProviderIdentity + endpoint_tag: str | None = None + endpoint_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 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) + + +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. + + 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: + """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) + if forwarded: + return forwarded + return public_model_id(getattr(model, "id")) + + +def public_model_id(model_id: str) -> str: + """Model id exposed by model-path API responses. + + 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.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. + + 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: + return canonical + 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 + + +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[EndpointIdentity] | None] = {} + self.rate_limited = False + + +async def _fetch_openrouter_endpoint_subproviders( + client: httpx.AsyncClient, + base_url: str, + api_key: str, + author_slug: str, + semaphore: asyncio.Semaphore, + cycle: _RefreshCycleState, +) -> 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 + "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[EndpointIdentity] | None + 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)}, + ) + cycle.endpoint_cache[cache_key] = None + return None + + if resp.status_code == 429: + logger.warning( + "OpenRouter endpoint discovery rate-limited; aborting cycle", + extra={"author_slug": author_slug}, + ) + 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}, + ) + cycle.endpoint_cache[cache_key] = None + return None + + try: + payload = resp.json() + data = payload.get("data") if isinstance(payload, dict) else None + endpoints = data.get("endpoints") if isinstance(data, dict) else None + if not isinstance(endpoints, list): + raise ValueError("endpoints must be a list") + identities: dict[str, EndpointIdentity] = {} + for endpoint in endpoints: + 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 endpoints and not identities: + raise ValueError("endpoints contain no usable tags") + result = list(identities.values()) + except Exception as e: # noqa: BLE001 + logger.warning( + "OpenRouter endpoint discovery bad payload", + extra={"author_slug": author_slug, "error": str(e)}, + ) + result = None + + cycle.endpoint_cache[cache_key] = result + return result + + +async def _load_model_visibility() -> tuple[ + dict[ModelKey, ModelRow], + set[ModelKey], + dict[int, ConfiguredProviderIdentity], +]: + """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 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( + selectinload(UpstreamProviderRow.models) # type: ignore[arg-type] + ) + provider_rows = (await session.exec(query)).all() + + overrides_by_key: dict[ModelKey, ModelRow] = {} + disabled_model_keys: set[ModelKey] = set() + provider_identities: dict[int, ConfiguredProviderIdentity] = {} + + for provider in provider_rows: + if not provider.enabled or provider.id is None: + continue + provider_identities[provider.id] = ConfiguredProviderIdentity( + 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) + if model.enabled: + overrides_by_key[key] = model + else: + disabled_model_keys.add(key) + + return overrides_by_key, disabled_model_keys, provider_identities + + +def _apply_model_visibility( + upstream: BaseUpstreamProvider, + 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. + + 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", "") + key = (model_id.lower(), upstream_provider_id) + if not getattr(model, "enabled", True) or key in disabled_model_keys: + continue + # 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()) + + # 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, + provider_identity: ConfiguredProviderIdentity, + overrides_by_key: dict[ModelKey, ModelRow] | None = None, + disabled_model_keys: set[ModelKey] | None = None, + cycle: _RefreshCycleState | None = None, +) -> ProviderPathSnapshot: + """Collect selectable routes while marking model-level degraded fetches. + + 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) + + def _base_path(model: object) -> DiscoveredPath: + model_id = exposed_model_id(model) + return DiscoveredPath( + model_id=model_id, + path=encode_model_path( + provider_identity.base_url, provider_identity.id, model_id + ), + provider=provider_identity, + ) + + if not is_openrouter_base_url(upstream.base_url): + return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models)) + + if not (upstream.provider_type or "").strip(): + return ProviderPathSnapshot(paths=()) + + semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) + async with _make_http_client() as client: + + async def _for_model( + model: object, + ) -> tuple[list[DiscoveredPath], str | None]: + model_id = exposed_model_id(model) + author_slug = openrouter_author_slug(model) + if not author_slug: + return [_base_path(model)], None + endpoints = await _fetch_openrouter_endpoint_subproviders( + client, + upstream.base_url, + upstream.api_key, + author_slug, + semaphore, + cycle, + ) + if endpoints is None: + return [], model_id + paths = [_base_path(model)] + paths.extend( + DiscoveredPath( + model_id=model_id, + 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, + ) + for endpoint in endpoints + ) + return paths, None + + results = await asyncio.gather( + *(_for_model(model) for model in models), return_exceptions=True + ) + + 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 + model_paths, preserved_model_id = result + paths.extend(model_paths) + if preserved_model_id: + preserve_model_ids.add(preserved_model_id) + + return ProviderPathSnapshot( + paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids) + ) + + +async def _persist_provider_paths( + upstream_provider_id: int, snapshot: ProviderPathSnapshot +) -> None: + """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: + delete_stmt = delete(ModelPathRow).where( + col(ModelPathRow.upstream_provider_id) == upstream_provider_id + ) + 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] + 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_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() + + +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 enabled_ids: + stmt = stmt.where( + 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() + + +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. + 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_key, + disabled_model_keys, + 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 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, + ) + if snapshot.preserve_model_ids: + logger.warning( + "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), + }, + ) + 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", + 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_for_provider(upstream_provider_id: int) -> None: + """Synchronize one provider when model-path discovery is enabled.""" + if _refresh_interval_seconds() <= 0: + return + + 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() + + +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 + + 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``. + + 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): + return upstreams_provider() + 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: + 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 + + +def _serialize_path(row: ModelPathRow) -> dict[str, Any]: + endpoint = None + if row.endpoint_tag or row.endpoint_name: + endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name} + return { + "path": row.path, + "provider": { + "id": row.upstream_provider_id, + "slug": row.provider_slug, + "type": row.provider_type, + }, + "endpoint": endpoint, + } + + +async def get_all_model_paths() -> dict: + """All models with their exact selectable routes.""" + async with create_session() as session: + rows = ( + 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[str, Any]]] = {} + seen_paths: dict[str, set[str]] = {} + updated_at = 0 + for row in rows: + updated_at = max(updated_at, row.updated_at) + if row.path in seen_paths.setdefault(row.model_id, set()): + continue + seen_paths[row.model_id].add(row.path) + grouped.setdefault(row.model_id, []).append(_serialize_path(row)) + data = [ + { + "id": grouped_model_id, + "paths": grouped[grouped_model_id], + } + 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: + """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() + ) + + 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] = [] + updated_at = 0 + for row in rows: + updated_at = max(updated_at, row.updated_at) + if row.path in seen: + continue + seen.add(row.path) + paths.append(_serialize_path(row)) + return {"data": paths, "updated_at": updated_at or None} diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 0cf78b18..c70fbcb9 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") diff --git a/tests/unit/test_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index e253412f..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,7 +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() - assert version == ("bf76270b66c4",) + 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 new file mode 100644 index 00000000..240a4199 --- /dev/null +++ b/tests/unit/test_model_paths.py @@ -0,0 +1,1475 @@ +"""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 + +import asyncio +import json +import os +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Any, AsyncGenerator, Callable, cast + +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 + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402 +from routstr.payment.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 +# --------------------------------------------------------------------------- # + + +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, + ) + + +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(BaseUpstreamProvider): + """Real ``BaseUpstreamProvider`` so the discovery-path hooks are the + production ones, with cached models injected.""" + + def __init__( + self, + *, + provider_type: str, + base_url: str, + models: list[SimpleNamespace], + db_id: int | None = 1, + api_key: str = "sk-test", + ) -> None: + 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]: # type: ignore[override] + return self._models + + +class _FakeOpenRouterProvider(OpenRouterUpstreamProvider): + """Real OpenRouter provider so the ``unknown`` mapping is the production one.""" + + 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( + *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) + + +@pytest.fixture +async def patched_session( + monkeypatch: pytest.MonkeyPatch, +) -> AsyncGenerator[AsyncEngine, None]: + """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 every provider id the tests insert path rows for. + async with AsyncSession(engine) as session: + for pid in _SEEDED_PROVIDER_IDS: + 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() + + +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"]} + + +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, + 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": _expected_path(provider_id, model_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 +# --------------------------------------------------------------------------- # + + +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_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" + ) + + +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_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_prefers_canonical() -> 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_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 + + +# --------------------------------------------------------------------------- # +# Refresh through the public entry point +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +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, + ) + 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, "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, +) -> 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, + ) + 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_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, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("enabled-model"), _model("disabled-model", enabled=False)], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"enabled-model"} + + +@pytest.mark.asyncio +async def test_disabling_model_on_one_provider_keeps_other_provider( + patched_session: AsyncEngine, +) -> None: + """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() + + 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") == {_expected_path(1, "shared-model")} + + +@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") == {_expected_path(1, "shared-model")} + assert _paths_of(payload, "private-alias") == {_expected_path(2, "private-alias")} + + +@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 +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]) + + assert (await mp.get_all_model_paths())["data"] == [] + + +@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]) + + assert (await mp.get_all_model_paths())["data"] == [ + {"id": "public-alias", "paths": [_path_entry(1, "public-alias")]} + ] + + +@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]) + + assert (await mp.get_all_model_paths())["data"] == [ + { + "id": "public-deployment", + "paths": [_path_entry(1, "public-deployment")], + } + ] + + +@pytest.mark.asyncio +async def test_refresh_replaces_stale_rows( + patched_session: AsyncEngine, +) -> None: + 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]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + await mp.refresh_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: + 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( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("fresh-model")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + + assert (await mp.get_all_model_paths())["data"] == [] + + +@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]) + assert (await mp.get_all_model_paths())["data"] == [] + + +@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="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) -> Any: + if upstream is bad: + raise RuntimeError("boom") + return await original(upstream, *args, **kwargs) + + monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) + + await mp.refresh_model_paths([good, bad]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} + + +# --------------------------------------------------------------------------- # +# 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( + ("Google", "google-vertex/eu"), + ("Google", "google-vertex/us"), + ), + ) + + await mp.refresh_model_paths([provider]) + + payload = await mp.get_paths_for_model("claude-opus-4.6") + assert {item["path"] for item in payload["data"]} == { + _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"] + } == {"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_uses_exact_tag_even_when_display_name_is_router( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """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, + ) + _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 paths == { + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "openrouter"), + } + + +@pytest.mark.asyncio +async def test_generic_provider_with_openrouter_base_url_discovers( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Configured provider identity is independent from its endpoint URL.""" + 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 == { + _expected_path(1, "claude-opus-4.6"), + _expected_path(1, "claude-opus-4.6", "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 +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") == { + _expected_path(2, "shared"), + _expected_path(2, "shared", "anthropic"), + _expected_path(2, "shared", "google"), + } + + +@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 _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) + + _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]) + expected = { + _expected_path(2, "m0"), + _expected_path(2, "m0", "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]) + + # 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") == expected + + +@pytest.mark.asyncio +async def test_openrouter_bad_payload_shapes_preserve_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """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}, + {}, + ): + + def _handler( + request: httpx.Request, p: dict[str, Any] | None = payload + ) -> httpx.Response: + return httpx.Response(200, json=p) + + _mock_transport(monkeypatch, _handler) + await mp.refresh_model_paths([provider]) + assert _paths_of(await mp.get_all_model_paths(), "m") == before + + +@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 paths == { + _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"), + } + + +@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"]} == { + _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"]) + + +@pytest.mark.asyncio +async def test_get_all_model_paths_keeps_distinct_configured_providers( + 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_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ], + } + ] + + +@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_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 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"] == [] + + +@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_entry(4, "glm-5v-turbo") + ] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_accepts_provider_prefixed_alias( + 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 short_paths == [ + _path_entry(4, "deepseek-v4-pro"), + _path_entry(7, "deepseek-v4-pro"), + ] + 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: + """Three-segment upstream IDs resolve to the same first-slash public ID.""" + 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_entry(1, "fireworks/models/glm-5") + ] + assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ + "data" + ] == [_path_entry(1, "fireworks/models/glm-5")] + + +# --------------------------------------------------------------------------- # +# 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_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, +) -> 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) + + 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" + + +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 +# --------------------------------------------------------------------------- # + + +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: + expected = { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + _path_entry(1, "claude-opus-4.6"), + _path_entry( + 2, + "claude-opus-4.6", + endpoint_tag="google-vertex/us", + endpoint_name="Google", + ), + ], + } + ], + "updated_at": 1753500000, + } + + async def _fake_get_all_model_paths() -> dict[str, Any]: + 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() == 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( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[str] = [] + + expected = { + "data": [ + _path_entry( + 2, + "anthropic/claude-opus-4.6", + 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 expected + + 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() == expected + assert calls == ["anthropic/claude-opus-4.6"] 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