mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
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:
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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}`
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
@@ -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}
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
+767
-336
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user