fix: address model-paths review findings

Provider scoping (items 1/2/6):
- Key visibility maps on (model_id.lower(), upstream_provider_id), matching
  refresh_model_maps, so a disable/override row on one provider never leaks
  onto another provider's model, and matching is case-insensitive.

Data safety (items 3/5):
- Degraded OpenRouter fetches (network error, 429, non-200, bad payload)
  return None (unknown) instead of []; a provider whose path set is unknown
  keeps its previously persisted rows instead of being wiped.
- Endpoint payload parsing moved fully inside try, with a list guard, so
  endpoints:null or non-list shapes are swallowed as documented.
- refresh with an empty live upstream list is a no-op; the unfiltered
  DELETE in the prune path is gone (prune now keys off enabled DB rows).

Hot path (items 4/12/14):
- Persist uses chunked bulk INSERTs (one statement per 500 rows) instead of
  per-row ORM adds; redundant ix_model_paths_model_id index dropped.
- Read routes filter in SQL instead of materializing the whole table, and
  output ordering is deterministic (public id + path), independent of rowid.
- Visibility no longer rebuilds fully priced Model objects per override row;
  it reads id/forwarded_model_id/canonical_slug straight off ModelRow.

Path/id contract (items 7/8/9/11):
- discovery_path_for_subprovider/discovery_base_paths hooks on
  BaseUpstreamProvider, overridden by OpenRouterUpstreamProvider, mirror
  _apply_provider_field so discovery and response stamping cannot drift
  (openrouter:OpenRouter now correctly maps to unknown).
- openrouter_author_slug falls back to a slash-containing forwarded_model_id,
  so admin-created alias rows are discoverable.
- public_model_id splits on the first slash, same as get_base_model_id, so
  discovery ids can be sent to chat completions verbatim.

Lifecycle (items 10/13):
- ENABLE_MODEL_PATHS_REFRESH kill switch; interval and flag re-read every
  loop iteration, and the task idles (not exits) while disabled.
- First 429 latches and aborts the remaining fan-out for the cycle; a
  per-cycle cache dedupes fetches across providers sharing a base URL.
- refresh_model_maps prunes paths of disabled/deleted providers so admin
  mutations take effect immediately; rows carry updated_at and both
  endpoints expose it.

Tests (item 15) rewritten through the public refresh entry point with
transport-level httpx.MockTransport fakes, FK enforcement on, and coverage
for the periodic loop. Migration re-chained onto 9c4d8e2f1a6b.
This commit is contained in:
9qeklajc
2026-07-26 13:23:31 +02:00
parent 0a00527626
commit f96acbb99c
13 changed files with 1183 additions and 552 deletions
+7 -4
View File
@@ -374,10 +374,13 @@ GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4
}
```
Model IDs in responses are unqualified display IDs: provider prefixes such as
`z-ai/` or `openai/` are stripped. Path values match the provider string stamped
on chat-completion responses, such as `anthropic`, `generic:my-upstream`, or
`openrouter:Anthropic`.
Model IDs in responses are base model IDs: the leading provider prefix such as
`z-ai/` or `openai/` is stripped (the same rule routing uses, so the ID can be
sent back to `/v1/chat/completions` verbatim). Path values match the provider
string stamped on chat-completion responses, such as `anthropic`,
`generic:Anthropic`, `openrouter:Anthropic`, or `unknown` (native OpenRouter
with no usable sub-provider). Responses also carry an `updated_at` Unix
timestamp of the last successful refresh (`null` when no refresh has run).
## Wallet Management
+2 -1
View File
@@ -142,7 +142,8 @@ Use environment variables for:
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
| `CORS_ORIGINS` | Allowed CORS origins | `*` |
| `RELAYS` | Nostr relays (comma-separated) | (default set) |
| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to disable | `600` |
| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to pause the refresh (previously discovered paths keep being served) | `600` |
| `ENABLE_MODEL_PATHS_REFRESH` | Kill switch for the background model-path refresh (OpenRouter endpoint fan-out) | `true` |
### Priority
@@ -1,7 +1,7 @@
"""add model paths table
Revision ID: 4e0c3d195a49
Revises: 7f2843d3f4e4
Revises: 9c4d8e2f1a6b
Create Date: 2026-07-24 21:14:39.062179
"""
@@ -11,19 +11,19 @@ from alembic import op
# revision identifiers, used by Alembic.
revision = "4e0c3d195a49"
down_revision = "7f2843d3f4e4"
down_revision = "9c4d8e2f1a6b"
branch_labels = None
depends_on = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.create_table(
"model_paths",
sa.Column("id", sa.Integer(), nullable=False),
sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("upstream_provider_id", sa.Integer(), nullable=False),
sa.Column("updated_at", sa.Integer(), nullable=False, server_default="0"),
sa.ForeignKeyConstraint(
["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE"
),
@@ -35,21 +35,16 @@ def upgrade() -> None:
name="uq_model_paths_model_path_provider",
),
)
op.create_index(
op.f("ix_model_paths_model_id"), "model_paths", ["model_id"], unique=False
)
# No standalone index on model_id: the unique constraint's autoindex already
# leads on model_id.
op.create_index(
op.f("ix_model_paths_upstream_provider_id"),
"model_paths",
["upstream_provider_id"],
unique=False,
)
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths")
op.drop_index(op.f("ix_model_paths_model_id"), table_name="model_paths")
op.drop_table("model_paths")
# ### end Alembic commands ###
+14 -11
View File
@@ -222,7 +222,10 @@ async def release_stale_reservations(
if released:
logger.warning(
"Released stale reservations",
extra={"released_reservations": released, "max_age_seconds": max_age_seconds},
extra={
"released_reservations": released,
"max_age_seconds": max_age_seconds,
},
)
return released
@@ -255,9 +258,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
.where(col(ApiKey.total_spent) == 0)
.where(col(ApiKey.total_requests) == 0)
.where(col(ApiKey.parent_key_hash).is_(None))
.where(
(col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)
)
.where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff))
.where(~pending_invoice)
.where(~has_children)
)
@@ -331,8 +332,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore
),
)
id: int | None = Field(default=None, primary_key=True)
# No standalone index on model_id: the unique constraint's autoindex already
# leads on model_id, so a second index only adds write amplification.
model_id: str = Field(
index=True, description="Client-visible /v1/models id (forwarded_model_id or id)"
description="Client-visible /v1/models id (forwarded_model_id or id)"
)
path: str = Field(
description="Provider path stamped on chat completion responses, e.g. "
@@ -344,6 +347,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore
ondelete="CASCADE",
description="upstream_providers.id this path was discovered from",
)
updated_at: int = Field(
default=0,
description="Unix timestamp of the refresh cycle that wrote this row",
)
class LightningInvoice(SQLModel, table=True): # type: ignore
@@ -631,9 +638,7 @@ class CliToken(SQLModel, table=True): # type: ignore
"""Long-lived authorization token for CLI/agent use against admin endpoints."""
__tablename__ = "cli_tokens"
id: str = Field(
primary_key=True, default_factory=lambda: uuid.uuid4().hex
)
id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex)
token: str = Field(unique=True, index=True, description="Bearer token value")
name: str = Field(description="Human-readable label for this token")
created_at: int = Field(default_factory=lambda: int(time.time()))
@@ -732,9 +737,7 @@ async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool:
return result.rowcount == 1
async def complete_routstr_fee_payout(
session: AsyncSession, paid_msats: int
) -> bool:
async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) -> bool:
"""Mark a checkpointed payout complete after the external payment succeeds."""
stmt = (
update(RoutstrFee)
+9 -14
View File
@@ -131,12 +131,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
refresh_upstreams_models_periodically(get_upstreams)
)
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
if global_settings.model_paths_refresh_interval_seconds > 0:
from ..upstream.model_paths import refresh_model_paths_periodically
# Always started: the loop re-reads the enable flag and interval every
# iteration, so 0 -> N (or re-enabling) takes effect without a restart.
from ..upstream.model_paths import refresh_model_paths_periodically
model_paths_refresh_task = asyncio.create_task(
refresh_model_paths_periodically(get_upstreams)
)
model_paths_refresh_task = asyncio.create_task(
refresh_model_paths_periodically(get_upstreams)
)
payout_task = asyncio.create_task(periodic_payout())
if global_settings.nsec:
nip91_task = asyncio.create_task(announce_provider())
@@ -144,9 +145,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
if global_settings.providers_refresh_interval_seconds > 0:
providers_task = asyncio.create_task(providers_cache_refresher())
key_reset_task = asyncio.create_task(periodic_key_reset())
stale_reservation_task = asyncio.create_task(
periodic_stale_reservation_sweep()
)
stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep())
dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune())
auto_topup_task = asyncio.create_task(periodic_auto_topup())
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
@@ -256,9 +255,7 @@ class _ImmutableStaticFiles(StaticFiles):
async def get_response(self, path: str, scope: Scope) -> StarletteResponse:
response = await super().get_response(path, scope)
if response.status_code == 200:
response.headers["Cache-Control"] = (
"public, max-age=31536000, immutable"
)
response.headers["Cache-Control"] = "public, max-age=31536000, immutable"
return response
@@ -332,9 +329,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
# Serve the App Router RSC payload for the home page.
@app.get("/index.txt", include_in_schema=False)
async def serve_root_rsc() -> FileResponse:
return FileResponse(
UI_DIST_PATH / "index.txt", media_type="text/x-component"
)
return FileResponse(UI_DIST_PATH / "index.txt", media_type="text/x-component")
# Next.js is built with `trailingSlash: true`, so all UI page URLs end
# with a slash (e.g. `/login/`). The proxy router catches `/{path:path}`
+13 -5
View File
@@ -100,8 +100,13 @@ class Settings(BaseSettings):
)
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
enable_model_paths_refresh: bool = Field(
default=True, env="ENABLE_MODEL_PATHS_REFRESH"
)
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(
default=604800, env="REFUND_SWEEP_TTL_SECONDS"
)
# Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL")
@@ -120,9 +125,8 @@ class Settings(BaseSettings):
# Discovery
relays: list[str] = Field(default_factory=list, env="RELAYS")
enable_analytics_sharing: bool = Field(
default=True, env="ENABLE_ANALYTICS_SHARING"
)
enable_analytics_sharing: bool = Field(default=True, env="ENABLE_ANALYTICS_SHARING")
def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
"""Discard unknown keys from persisted settings."""
@@ -330,7 +334,11 @@ class SettingsService:
valid_fields = set(env_resolved.dict().keys())
merged_dict: dict[str, Any] = dict(env_resolved.dict())
merged_dict.update(
{k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields}
{
k: v
for k, v in db_json.items()
if v not in (None, "", [], {}) and k in valid_fields
}
)
merged_dict = Settings(**merged_dict).dict()
+6 -6
View File
@@ -455,7 +455,9 @@ async def _update_sats_pricing_once() -> None:
for m in upstream.get_cached_models()
]
upstream._models_cache = updated_models
upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models}
upstream._models_by_id = {
m.forwarded_model_id or m.id: m for m in updated_models
}
updated_count += len(updated_models)
if updated_count > 0:
@@ -510,9 +512,7 @@ class ModelTestRequest(V2BaseModel):
request_data: dict
@models_router.post(
"/api/models/test", dependencies=[Depends(_require_admin_api)]
)
@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)])
async def test_model(
payload: ModelTestRequest,
session: AsyncSession = Depends(get_session),
@@ -601,7 +601,7 @@ async def model_paths() -> dict:
"""All models with every upstream provider path they are reachable through."""
from ..upstream.model_paths import get_all_model_paths
return {"data": await get_all_model_paths()}
return await get_all_model_paths()
@models_router.get("/v1/models/paths/model")
@@ -615,7 +615,7 @@ async def model_paths_for_model(model_id: str) -> dict:
"""
from ..upstream.model_paths import get_paths_for_model
return {"data": await get_paths_for_model(model_id)}
return await get_paths_for_model(model_id)
@models_router.get("/v1/models")
+13
View File
@@ -183,6 +183,19 @@ async def refresh_model_maps() -> None:
disabled_model_keys=disabled_model_keys,
)
# Keep model-path discovery in sync with admin mutations: disabling or
# deleting a provider must stop advertising its paths immediately rather
# than after the next timed refresh.
from .upstream.model_paths import prune_model_paths_for_inactive_providers
try:
await prune_model_paths_for_inactive_providers()
except Exception as e: # noqa: BLE001 - discovery sync must not break routing
logger.warning(
"Failed to prune model paths for inactive providers",
extra={"error": str(e), "error_type": type(e).__name__},
)
async def refresh_model_maps_periodically() -> None:
"""Background task to refresh model maps every minute."""
+22
View File
@@ -237,6 +237,28 @@ class BaseUpstreamProvider:
except (TypeError, ValueError):
pass
def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None:
"""Discovery path for a reported sub-provider name.
Must produce exactly the value ``_apply_provider_field`` would stamp on
a response whose upstream payload reported ``sub_provider``, so the
model-path discovery API never advertises a path that cannot appear on
the wire. Subclasses that override ``_apply_provider_field`` must
override this to match.
"""
provider_type = (self.provider_type or "").strip()
if not provider_type:
return None
sub = (sub_provider or "").strip()
if not sub or sub == provider_type or sub.startswith(f"{provider_type}:"):
return sub or provider_type
return f"{provider_type}:{sub}"
def discovery_base_paths(self) -> list[str]:
"""Paths stamped when the upstream reports no sub-provider of its own."""
provider_type = (self.provider_type or "").strip()
return [provider_type] if provider_type else []
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp the routstr ``provider`` field onto an upstream response payload.
+305 -164
View File
@@ -5,32 +5,34 @@ This is discovery/visibility data only — routing still selects the cheapest or
best provider separately.
A *path* is the provider string that may appear in Routstr chat completion
responses (see ``BaseUpstreamProvider._apply_provider_field``):
responses. The strings emitted here are produced by the provider's own
``discovery_path_for_subprovider`` / ``discovery_base_paths`` hooks, which
mirror ``_apply_provider_field`` so discovery and response stamping cannot
drift:
- Direct upstream -> ``<provider_type>`` e.g. ``anthropic``
- Generic/custom OpenRouter-compatible upstream -> ``generic:<name>``
- Native OpenRouter routing to a sub-provider -> ``openrouter:<name>``
Native OpenRouter does not emit a useful bare ``openrouter`` path when no
sub-provider is present; it reports ``unknown`` instead.
- Native OpenRouter with no usable sub-provider -> ``unknown``
"""
from __future__ import annotations
import asyncio
import random
import time
from typing import TYPE_CHECKING, Callable
import httpx
from sqlalchemy import insert, or_
from sqlalchemy.orm import selectinload
from sqlmodel import col, delete, select
from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session
from ..core.logging import get_logger
from .base import BaseUpstreamProvider
if TYPE_CHECKING:
from ..payment.models import Model
from .base import BaseUpstreamProvider
logger = get_logger(__name__)
@@ -39,6 +41,20 @@ logger = get_logger(__name__)
_OPENROUTER_CONCURRENCY = 5
_OPENROUTER_TIMEOUT_SECONDS = 10.0
# Rows inserted per statement during persist. Keeps each INSERT bounded while
# avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle.
_PERSIST_CHUNK_SIZE = 500
# Visibility key used across this module: routing carries the provider
# dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)),
# so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps.
ModelKey = tuple[str, int]
def _make_http_client() -> httpx.AsyncClient:
"""Client factory, separated so tests can substitute a mock transport."""
return httpx.AsyncClient()
def is_openrouter_base_url(base_url: str | None) -> bool:
"""True when ``base_url`` points at OpenRouter.
@@ -61,19 +77,21 @@ def exposed_model_id(model: object) -> str:
def public_model_id(model_id: str) -> str:
"""Model id exposed by model-path API responses.
Provider-prefixed ids such as ``z-ai/glm-5v-turbo`` are returned as
``glm-5v-turbo`` so clients can search and display the same unqualified id
they pass to ``/v1/models/paths/model``.
Uses the same rule as ``create_model_mappings.get_base_model_id`` and
``resolve_model_alias`` — strip everything before the *first* slash — so
the id shown here can be sent back to ``/v1/chat/completions`` verbatim.
"""
return model_id.rsplit("/", 1)[-1]
return model_id.split("/", 1)[1] if "/" in model_id else model_id
def openrouter_author_slug(model: object) -> str | None:
"""Return a canonical ``author/slug`` for the OpenRouter endpoints API.
OpenRouter requires the canonical id, never ``forwarded_model_id``. Prefer
``canonical_slug``, then a slash-containing ``id``; otherwise there is no
usable form and endpoint discovery is skipped for this model.
Prefer ``canonical_slug``, then a slash-containing ``id``, then a
slash-containing ``forwarded_model_id``. The forwarded id is exactly what
the proxy sends upstream for admin-created alias rows (``base.py`` forwards
``forwarded_model_id or id``), so it is a valid OpenRouter id when the
bare ``id`` is a local alias with no slash.
"""
canonical = getattr(model, "canonical_slug", None)
if canonical and "/" in canonical:
@@ -81,24 +99,51 @@ def openrouter_author_slug(model: object) -> str | None:
model_id = getattr(model, "id", None)
if model_id and "/" in model_id:
return model_id
forwarded = getattr(model, "forwarded_model_id", None)
if forwarded and "/" in forwarded:
return forwarded
return None
async def _fetch_openrouter_endpoint_paths(
class _RefreshCycleState:
"""Per-refresh shared state: fetch dedupe cache and rate-limit latch.
``endpoint_cache`` dedupes byte-identical ``/endpoints`` fetches when two
providers point at the same OpenRouter base URL. ``rate_limited`` latches
on the first 429 so the rest of the cycle stops hammering a throttled API;
the whole provider result then degrades to "unknown" instead of an empty
list, which preserves previously persisted rows.
"""
def __init__(self) -> None:
self.endpoint_cache: dict[tuple[str, str], list[str] | None] = {}
self.rate_limited = False
async def _fetch_openrouter_endpoint_subproviders(
client: httpx.AsyncClient,
base_url: str,
api_key: str,
author_slug: str,
path_prefix: str,
semaphore: asyncio.Semaphore,
) -> list[str]:
"""Return ``<path_prefix>:<provider_name>`` paths for one model, or ``[]``.
cycle: _RefreshCycleState,
) -> list[str] | None:
"""Return sub-provider names for one model, or ``None`` when unknown.
Failures (network, rate limit, bad payload) are logged and swallowed so one
model never breaks the whole refresh.
``None`` (not ``[]``) signals a degraded fetch — network failure, rate
limit, non-200, or an unparseable payload — so callers can distinguish
"this model has no endpoints" from "we could not find out". Failures are
logged and swallowed so one model never breaks the whole refresh.
"""
cache_key = (base_url, author_slug)
if cache_key in cycle.endpoint_cache:
return cycle.endpoint_cache[cache_key]
if cycle.rate_limited:
return None
url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints"
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
result: list[str] | None
async with semaphore:
try:
resp = await client.get(
@@ -109,48 +154,60 @@ async def _fetch_openrouter_endpoint_paths(
"OpenRouter endpoint discovery request failed",
extra={"author_slug": author_slug, "error": str(e)},
)
return []
cycle.endpoint_cache[cache_key] = None
return None
if resp.status_code == 429:
logger.warning(
"OpenRouter endpoint discovery rate-limited",
"OpenRouter endpoint discovery rate-limited; aborting cycle",
extra={"author_slug": author_slug},
)
return []
cycle.rate_limited = True
cycle.endpoint_cache[cache_key] = None
return None
if resp.status_code != 200:
logger.warning(
"OpenRouter endpoint discovery non-200",
extra={"author_slug": author_slug, "status_code": resp.status_code},
)
return []
cycle.endpoint_cache[cache_key] = None
return None
try:
endpoints = resp.json().get("data", {}).get("endpoints", [])
if not isinstance(endpoints, list):
endpoints = []
names: list[str] = []
for endpoint in endpoints:
provider_name = (
endpoint.get("provider_name") if isinstance(endpoint, dict) else None
)
if provider_name:
names.append(provider_name)
result = list(dict.fromkeys(names))
except Exception as e: # noqa: BLE001
logger.warning(
"OpenRouter endpoint discovery bad payload",
extra={"author_slug": author_slug, "error": str(e)},
)
return []
result = None
paths: list[str] = []
for endpoint in endpoints:
provider_name = (endpoint or {}).get("provider_name")
if provider_name:
paths.append(f"{path_prefix}:{provider_name}")
# De-duplicate while preserving order.
return list(dict.fromkeys(paths))
cycle.endpoint_cache[cache_key] = result
return result
async def _load_model_visibility() -> tuple[
dict[str, tuple[ModelRow, float]], set[str], set[int]
dict[ModelKey, ModelRow], set[ModelKey], set[int]
]:
"""Load the same DB model visibility inputs used by routing.
``refresh_model_maps`` builds routing from enabled providers, enabled DB
override rows, and disabled model ids. Model-path discovery uses the same
view so the discovery API does not advertise models routing would hide and
reports forwarded aliases from DB overrides consistently with ``/v1/models``.
override rows, and disabled model keys — all keyed on
``(model_id.lower(), upstream_provider_id)`` because ``ModelRow``'s primary
key is composite and the same id legitimately exists on several providers.
Model-path discovery uses the same keying so disabling a model on one
provider never hides it on another, and one provider's
``forwarded_model_id`` alias is never applied to a different provider.
"""
async with create_session() as session:
query = select(UpstreamProviderRow).options(
@@ -158,149 +215,149 @@ async def _load_model_visibility() -> tuple[
)
provider_rows = (await session.exec(query)).all()
overrides_by_id: dict[str, tuple[ModelRow, float]] = {}
disabled_model_ids: set[str] = set()
overrides_by_key: dict[ModelKey, ModelRow] = {}
disabled_model_keys: set[ModelKey] = set()
enabled_provider_ids: set[int] = set()
for provider in provider_rows:
if not provider.enabled:
if not provider.enabled or provider.id is None:
continue
if provider.id is not None:
enabled_provider_ids.add(provider.id)
enabled_provider_ids.add(provider.id)
for model in provider.models:
key = (model.id.lower(), provider.id)
if model.enabled:
overrides_by_id[model.id] = (model, provider.provider_fee)
overrides_by_key[key] = model
else:
disabled_model_ids.add(model.id)
disabled_model_keys.add(key)
return overrides_by_id, disabled_model_ids, enabled_provider_ids
def _row_to_visible_model(
model_id: str,
row: ModelRow,
provider_fee: float,
) -> Model | None:
"""Convert an enabled DB override row into a routed model object."""
from ..payment.models import _row_to_model
try:
return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee)
except Exception as exc: # noqa: BLE001 - skip invalid override row
logger.warning(
"Skipping invalid model override while collecting model paths",
extra={
"model_id": model_id,
"upstream_provider_id": getattr(row, "upstream_provider_id", None),
"error": str(exc),
"error_type": type(exc).__name__,
},
)
return None
return overrides_by_key, disabled_model_keys, enabled_provider_ids
def _apply_model_visibility(
upstream: BaseUpstreamProvider,
overrides_by_id: dict[str, tuple[ModelRow, float]] | None,
disabled_model_ids: set[str] | None,
overrides_by_key: dict[ModelKey, ModelRow] | None,
disabled_model_keys: set[ModelKey] | None,
) -> list[object]:
"""Return provider models after DB disabled/override state is applied."""
overrides_by_id = overrides_by_id or {}
disabled_model_ids = disabled_model_ids or set()
"""Return provider models after DB disabled/override state is applied.
Only the identity fields (``id``, ``forwarded_model_id``,
``canonical_slug``) matter for path discovery, so DB override rows are used
directly rather than rebuilt into fully priced ``Model`` objects — the
pricing pipeline costs ~0.7ms of event-loop CPU per row for data this
module immediately discards.
"""
overrides_by_key = overrides_by_key or {}
disabled_model_keys = disabled_model_keys or set()
upstream_provider_id = getattr(upstream, "db_id", None)
if not isinstance(upstream_provider_id, int):
return [
model
for model in upstream.get_cached_models()
if getattr(model, "enabled", True)
]
visible_models: list[object] = []
seen_model_ids: set[str] = set()
for model in upstream.get_cached_models():
model_id = getattr(model, "id", "")
if not getattr(model, "enabled", True) or model_id in disabled_model_ids:
key = (model_id.lower(), upstream_provider_id)
if not getattr(model, "enabled", True) or key in disabled_model_keys:
continue
if model_id in overrides_by_id:
override_row, provider_fee = overrides_by_id[model_id]
visible_model = _row_to_visible_model(model_id, override_row, provider_fee)
if visible_model is None:
continue
model = visible_model
if not getattr(model, "enabled", True):
continue
visible_models.append(model)
# Apply overrides only for this provider's own model row.
override_row = overrides_by_key.get(key)
visible: object = model if override_row is None else override_row
visible_models.append(visible)
seen_model_ids.add(model_id.lower())
upstream_provider_id = getattr(upstream, "db_id", None)
if isinstance(upstream_provider_id, int):
for model_id, (override_row, provider_fee) in overrides_by_id.items():
if model_id in disabled_model_ids:
continue
if (
getattr(override_row, "upstream_provider_id", None)
!= upstream_provider_id
):
continue
if model_id.lower() in seen_model_ids:
continue
override_model = _row_to_visible_model(model_id, override_row, provider_fee)
if override_model is None:
continue
if getattr(override_model, "enabled", True):
visible_models.append(override_model)
seen_model_ids.add(model_id.lower())
# DB-only override rows for this provider with no cached counterpart.
for (model_id_lower, provider_id), override_row in overrides_by_key.items():
if provider_id != upstream_provider_id:
continue
if model_id_lower in seen_model_ids:
continue
visible_models.append(override_row)
seen_model_ids.add(model_id_lower)
return visible_models
async def _collect_provider_paths(
upstream: BaseUpstreamProvider,
overrides_by_id: dict[str, tuple[ModelRow, float]] | None = None,
disabled_model_ids: set[str] | None = None,
) -> list[tuple[str, str]]:
overrides_by_key: dict[ModelKey, ModelRow] | None = None,
disabled_model_keys: set[ModelKey] | None = None,
cycle: _RefreshCycleState | None = None,
) -> list[tuple[str, str]] | None:
"""Collect ``(model_id, path)`` pairs for one provider instance.
Emits the direct ``<provider_type>`` path for normal upstreams. For
OpenRouter-compatible providers, emits one path per OpenRouter sub-provider
endpoint, prefixed the same way response stamping prefixes it.
Emits the provider's ``discovery_base_paths`` for normal upstreams. For
OpenRouter-compatible providers, additionally emits one path per OpenRouter
sub-provider endpoint via ``discovery_path_for_subprovider`` so the strings
match response stamping exactly.
Returns ``None`` when the provider's path set could not be determined this
cycle (every endpoint fetch degraded); callers must then keep previously
persisted rows instead of wiping them.
"""
provider_type = (upstream.provider_type or "").strip()
models = _apply_model_visibility(upstream, overrides_by_id, disabled_model_ids)
is_openrouter = is_openrouter_base_url(upstream.base_url)
cycle = cycle or _RefreshCycleState()
models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys)
base_paths = upstream.discovery_base_paths()
pairs: list[tuple[str, str]] = []
if not is_openrouter:
for model in models:
if provider_type:
pairs.append((exposed_model_id(model), provider_type))
return pairs
if not is_openrouter_base_url(upstream.base_url):
return [
(exposed_model_id(model), path) for model in models for path in base_paths
]
if not provider_type:
return pairs
if not (upstream.provider_type or "").strip():
return []
any_fetch_succeeded = False
any_fetch_attempted = False
semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY)
async with httpx.AsyncClient() as client:
async with _make_http_client() as client:
async def _for_model(model: object) -> list[tuple[str, str]]:
nonlocal any_fetch_succeeded, any_fetch_attempted
model_id = exposed_model_id(model)
# Base paths always apply: responses whose upstream payload lacks a
# provider field are stamped with them (see _apply_provider_field).
pairs = [(model_id, path) for path in base_paths]
author_slug = openrouter_author_slug(model)
if not author_slug:
return []
paths = await _fetch_openrouter_endpoint_paths(
return pairs
any_fetch_attempted = True
sub_providers = await _fetch_openrouter_endpoint_subproviders(
client,
upstream.base_url,
upstream.api_key,
author_slug,
provider_type,
semaphore,
cycle,
)
model_id = exposed_model_id(model)
return [(model_id, path) for path in paths]
if sub_providers is None:
return []
any_fetch_succeeded = True
paths = [
upstream.discovery_path_for_subprovider(name) for name in sub_providers
]
pairs.extend((model_id, path) for path in paths if path)
return list(dict.fromkeys(pairs))
results = await asyncio.gather(
*(_for_model(m) for m in models), return_exceptions=True
)
if any_fetch_attempted and not any_fetch_succeeded:
# Every endpoint lookup degraded (offline, throttled, bad payloads):
# the true path set is unknown, not empty.
return None
pairs: list[tuple[str, str]] = []
for result in results:
if isinstance(result, BaseException):
logger.warning(
"OpenRouter endpoint discovery task errored",
extra={"provider": provider_type, "error": str(result)},
extra={"provider": upstream.provider_type, "error": str(result)},
)
continue
pairs.extend(result)
@@ -314,33 +371,57 @@ async def _persist_provider_paths(
"""Replace all rows for ``upstream_provider_id`` with ``pairs``.
Replacement (not upsert) so stale paths disappear when provider config or
upstream availability changes.
upstream availability changes. Rows are written with chunked bulk INSERTs
so the transaction holds SQLite's write lock briefly — billing writes share
this database file.
"""
unique_pairs = list(dict.fromkeys(pairs))
now = int(time.time())
async with create_session() as session:
await session.exec( # type: ignore[call-overload]
delete(ModelPathRow).where(
col(ModelPathRow.upstream_provider_id) == upstream_provider_id
)
)
for model_id, path in unique_pairs:
session.add(
ModelPathRow(
model_id=model_id,
path=path,
upstream_provider_id=upstream_provider_id,
)
for start in range(0, len(unique_pairs), _PERSIST_CHUNK_SIZE):
chunk = unique_pairs[start : start + _PERSIST_CHUNK_SIZE]
await session.execute(
insert(ModelPathRow),
[
{
"model_id": model_id,
"path": path,
"upstream_provider_id": upstream_provider_id,
"updated_at": now,
}
for model_id, path in chunk
],
)
await session.commit()
async def _prune_inactive_provider_paths(active_provider_ids: set[int]) -> None:
"""Delete paths for providers no longer present in the live upstream set."""
async def prune_model_paths_for_inactive_providers() -> None:
"""Delete paths whose provider is no longer enabled in the database.
Called from ``refresh_model_maps`` so admin mutations (disable/delete
provider) stop advertising a provider's paths immediately instead of
waiting for the next timed refresh. Uses the DB as the source of truth, so
it is safe at boot even before upstreams initialize.
"""
async with create_session() as session:
enabled_ids = (
await session.exec(
select(UpstreamProviderRow.id).where(
col(UpstreamProviderRow.enabled).is_(True)
)
)
).all()
stmt = delete(ModelPathRow)
if active_provider_ids:
if enabled_ids:
stmt = stmt.where(
col(ModelPathRow.upstream_provider_id).not_in(active_provider_ids)
col(ModelPathRow.upstream_provider_id).not_in(
[pid for pid in enabled_ids if pid is not None]
)
)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
@@ -352,28 +433,42 @@ async def refresh_model_paths(
"""Recompute and persist model paths for every enabled provider.
One provider's failure is logged and isolated; it must not break the rest.
A provider whose paths could not be determined this cycle keeps its
previously persisted rows. An empty ``upstreams`` list (e.g. a failed
``initialize_upstreams`` at boot) is treated as "unknown" and touches
nothing.
"""
if not upstreams:
logger.warning("Skipping model paths refresh: no live upstreams")
return
(
overrides_by_id,
disabled_model_ids,
overrides_by_key,
disabled_model_keys,
enabled_provider_ids,
) = await _load_model_visibility()
active_provider_ids = {
upstream.db_id
for upstream in upstreams
if upstream.db_id is not None and upstream.db_id in enabled_provider_ids
}
await _prune_inactive_provider_paths(active_provider_ids)
await prune_model_paths_for_inactive_providers()
cycle = _RefreshCycleState()
for upstream in upstreams:
if upstream.db_id is None or upstream.db_id not in enabled_provider_ids:
continue
try:
pairs = await _collect_provider_paths(
upstream,
overrides_by_id=overrides_by_id,
disabled_model_ids=disabled_model_ids,
overrides_by_key=overrides_by_key,
disabled_model_keys=disabled_model_keys,
cycle=cycle,
)
if pairs is None:
logger.warning(
"Model paths unknown this cycle; keeping previous rows",
extra={
"provider": upstream.provider_type or upstream.base_url,
"db_id": upstream.db_id,
},
)
continue
await _persist_provider_paths(upstream.db_id, pairs)
except Exception as e: # noqa: BLE001 - isolate per-provider failures
logger.error(
@@ -387,18 +482,27 @@ async def refresh_model_paths(
)
def _refresh_interval_seconds() -> int:
"""Current interval, re-read every loop so runtime setting changes apply."""
from ..core.settings import settings
if not getattr(settings, "enable_model_paths_refresh", True):
return 0
return int(getattr(settings, "model_paths_refresh_interval_seconds", 0) or 0)
async def refresh_model_paths_periodically(
upstreams_provider: (
Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider]
),
) -> None:
"""Background task mirroring ``refresh_upstreams_models_periodically``."""
from ..core.settings import settings
"""Background task mirroring ``refresh_upstreams_models_periodically``.
interval = getattr(settings, "model_paths_refresh_interval_seconds", 0)
if not interval or interval <= 0:
logger.info("Model paths refresh disabled (interval <= 0)")
return
The interval and enable flag are re-read every iteration, so the refresh
can be turned off (or on) and retuned without a restart. While disabled the
task idles instead of exiting, so re-enabling takes effect.
"""
_DISABLED_POLL_SECONDS = 60.0
def _resolve_upstreams() -> list[BaseUpstreamProvider]:
if callable(upstreams_provider):
@@ -406,6 +510,14 @@ async def refresh_model_paths_periodically(
return upstreams_provider
while True:
interval = _refresh_interval_seconds()
if interval <= 0:
try:
await asyncio.sleep(_DISABLED_POLL_SECONDS)
except asyncio.CancelledError:
break
continue
try:
await refresh_model_paths(_resolve_upstreams())
except asyncio.CancelledError:
@@ -423,50 +535,79 @@ async def refresh_model_paths_periodically(
break
async def get_all_model_paths() -> list[dict]:
async def get_all_model_paths() -> dict:
"""All models with their paths, shaped for ``GET /v1/models/paths``."""
async with create_session() as session:
rows = (
await session.exec(select(ModelPathRow).order_by(ModelPathRow.model_id))
await session.exec(
select(ModelPathRow).order_by(
col(ModelPathRow.model_id),
col(ModelPathRow.path),
col(ModelPathRow.upstream_provider_id),
)
)
).all()
grouped: dict[str, list[dict]] = {}
seen_paths: dict[str, set[str]] = {}
updated_at = 0
for row in rows:
updated_at = max(updated_at, row.updated_at)
model_id = public_model_id(row.model_id)
if row.path in seen_paths.setdefault(model_id, set()):
continue
seen_paths[model_id].add(row.path)
grouped.setdefault(model_id, []).append({"path": row.path})
return [{"id": model_id, "paths": paths} for model_id, paths in grouped.items()]
# Deterministic output: models sorted by public id, paths sorted within.
data: list[dict] = []
for grouped_model_id in sorted(grouped):
model_paths = sorted(grouped[grouped_model_id], key=lambda p: str(p["path"]))
data.append({"id": grouped_model_id, "paths": model_paths})
return {"data": data, "updated_at": updated_at or None}
async def get_paths_for_model(model_id: str) -> list[dict]:
async def get_paths_for_model(model_id: str) -> dict:
"""Paths for a single model, shaped for ``GET /v1/models/paths/model``.
Match by the public, unqualified model id, mirroring the model cache alias
behavior. Both ``deepseek-v4-pro`` and ``deepseek/deepseek-v4-pro`` resolve
every row whose stored id has the same base model id.
every row whose stored id has the same base model id. The candidate set is
narrowed in SQL (exact id or ``%/<id>`` suffix) so the route does not
materialize the whole table per request.
"""
requested_id = public_model_id(model_id)
# The request may be a full stored id ("z-ai/glm-5v-turbo") or an
# already-stripped public id ("fireworks/models/glm-5"); accept both.
accepted_ids = {model_id, public_model_id(model_id)}
async with create_session() as session:
conditions = []
for candidate in accepted_ids:
conditions.append(col(ModelPathRow.model_id) == candidate)
conditions.append(col(ModelPathRow.model_id).endswith(f"/{candidate}"))
rows = (
await session.exec(
select(ModelPathRow).order_by(
ModelPathRow.path,
select(ModelPathRow)
.where(or_(*conditions))
.order_by(
col(ModelPathRow.path),
col(ModelPathRow.upstream_provider_id),
ModelPathRow.model_id,
col(ModelPathRow.model_id),
)
)
).all()
seen: set[str] = set()
paths: list[dict] = []
updated_at = 0
for row in rows:
if public_model_id(row.model_id) != requested_id:
# The SQL suffix match is a prefilter; enforce the exact public-id rule.
if (
row.model_id not in accepted_ids
and public_model_id(row.model_id) not in accepted_ids
):
continue
updated_at = max(updated_at, row.updated_at)
if row.path in seen:
continue
seen.add(row.path)
paths.append({"path": row.path})
return paths
return {"data": paths, "updated_at": updated_at or None}
+17
View File
@@ -18,6 +18,23 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
supports_anthropic_messages = True
litellm_provider_prefix = "openrouter/"
def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None:
"""Mirror ``_apply_provider_field``: strip repeated prefixes, map a
missing or self-echoing sub-provider to the literal ``"unknown"``."""
provider_type = (self.provider_type or "").strip()
sub = (sub_provider or "").strip()
prefix = f"{provider_type}:"
while sub.lower().startswith(prefix.lower()):
sub = sub[len(prefix) :].strip()
if not sub or sub.lower() == provider_type.lower():
return "unknown"
return f"{provider_type}:{sub}"
def discovery_base_paths(self) -> list[str]:
"""Native OpenRouter never stamps a bare ``openrouter``; a response
with no sub-provider is stamped ``unknown``."""
return ["unknown"]
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp the ``provider`` field for OpenRouter responses.
+3 -1
View File
@@ -37,7 +37,9 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None:
"payout_in_progress_msats, payout_started_at FROM routstr_fees"
).fetchone()
assert version == ("9c4d8e2f1a6b",)
# Head of the 7f2843d3f4e4 lineage: model-paths chains onto the fee-payout
# repair migration.
assert version == ("4e0c3d195a49",)
assert {
"id",
"accumulated_msats",
File diff suppressed because it is too large Load Diff