From e2f89a26454debb3b44860508209457eaff50ab6 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 27 Jul 2026 23:17:13 +0200 Subject: [PATCH] fix: address follow-up model path review --- docs/api/endpoints.md | 6 +- routstr/core/admin.py | 20 +--- routstr/core/db.py | 2 +- routstr/upstream/base.py | 22 ---- routstr/upstream/model_paths.py | 114 +++++++++++++++---- routstr/upstream/openrouter.py | 17 --- tests/unit/test_fee_payout_migration.py | 10 +- tests/unit/test_model_paths.py | 144 ++++++++++++++++++------ 8 files changed, 213 insertions(+), 122 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 2d202659..ca54b8e0 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -345,12 +345,12 @@ GET /v1/models/paths "id": "anthropic/claude-sonnet-4", "paths": [ { - "path": "provider=12", + "path": "provider=anthropic-primary", "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, "endpoint": null }, { - "path": "provider=42&endpoint=google-vertex%2Fus", + "path": "provider=openrouter-main&endpoint=google-vertex%2Fus", "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, "endpoint": {"tag": "google-vertex/us", "name": "Google"} } @@ -363,7 +363,7 @@ GET /v1/models/paths `path` is an opaque, percent-encoded selector. Clients must store and return it unchanged rather than parsing or reconstructing it. The configured provider's -stable node-local ID defines the upstream route; no upstream URL is exposed. +public slug defines the upstream route; no upstream URL is exposed. OpenRouter routes additionally use the exact machine-readable endpoint `tag`. Provider slugs/types and endpoint names are display data and never participate in identity. When request-side selection is implemented, an endpoint tag must diff --git a/routstr/core/admin.py b/routstr/core/admin.py index f581b704..0402f1a7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -52,20 +52,10 @@ MAX_USAGE_ANALYTICS_HOURS = 365 * 24 async def _refresh_provider_model_paths(upstream_provider_id: int) -> None: - """Best-effort immediate discovery sync after an admin mutation.""" - from ..upstream.model_paths import refresh_model_paths_for_provider + """Queue discovery sync without blocking the committed admin mutation.""" + from ..upstream.model_paths import schedule_model_paths_refresh_for_provider - try: - await refresh_model_paths_for_provider(upstream_provider_id) - except Exception as exc: # noqa: BLE001 - committed admin writes must survive - logger.warning( - "Failed to refresh model paths after admin mutation", - extra={ - "upstream_provider_id": upstream_provider_id, - "error": str(exc), - "error_type": type(exc).__name__, - }, - ) + await schedule_model_paths_refresh_for_provider(upstream_provider_id) async def require_admin_api(request: Request) -> None: @@ -725,9 +715,6 @@ async def batch_override_provider_models( json.dumps(model_data.alias_ids) if model_data.alias_ids else None ) existing_row.enabled = model_data.enabled - existing_row.forwarded_model_id = ( - model_data.forwarded_model_id or model_data.id - ) session.add(existing_row) else: # Create new @@ -758,7 +745,6 @@ async def batch_override_provider_models( ), upstream_provider_id=provider_pk, enabled=model_data.enabled, - forwarded_model_id=model_data.forwarded_model_id or model_data.id, ) session.add(row) diff --git a/routstr/core/db.py b/routstr/core/db.py index 151d8627..3edfe140 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -338,7 +338,7 @@ class ModelPathRow(SQLModel, table=True): # type: ignore description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( - description="Opaque selector containing provider ID and optional endpoint tag" + description="Opaque selector containing provider slug and optional endpoint tag" ) provider_slug: str = Field( description="Public slug of the configured upstream provider" diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 5f9116c9..a8dba7e3 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -237,28 +237,6 @@ class BaseUpstreamProvider: except (TypeError, ValueError): pass - def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: - """Discovery path for a reported sub-provider name. - - Must produce exactly the value ``_apply_provider_field`` would stamp on - a response whose upstream payload reported ``sub_provider``, so the - model-path discovery API never advertises a path that cannot appear on - the wire. Subclasses that override ``_apply_provider_field`` must - override this to match. - """ - provider_type = (self.provider_type or "").strip() - if not provider_type: - return None - sub = (sub_provider or "").strip() - if not sub or sub == provider_type or sub.startswith(f"{provider_type}:"): - return sub or provider_type - return f"{provider_type}:{sub}" - - def discovery_base_paths(self) -> list[str]: - """Paths stamped when the upstream reports no sub-provider of its own.""" - provider_type = (self.provider_type or "").strip() - return [provider_type] if provider_type else [] - def _apply_provider_field(self, response_json: object) -> None: """Stamp the routstr ``provider`` field onto an upstream response payload. diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 45ce6887..40022aca 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -5,12 +5,12 @@ This PR remains discovery-only: request-side routing will consume the opaque selectors in a follow-up. A path is a standard percent-encoded query string containing the configured -provider's stable node-local ID and, for an exact OpenRouter endpoint, its +provider's public slug and, for an exact OpenRouter endpoint, its machine-readable tag. Upstream URLs and display names never participate in or leak through public identity:: - provider=42 - provider=42&endpoint=google-vertex%2Fus + provider=anthropic-primary + provider=openrouter-main&endpoint=google-vertex%2Fus """ from __future__ import annotations @@ -23,7 +23,7 @@ from typing import TYPE_CHECKING, Any, Callable from urllib.parse import urlencode import httpx -from sqlalchemy import insert +from sqlalchemy.dialects.sqlite import insert from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select @@ -44,6 +44,12 @@ _OPENROUTER_TIMEOUT_SECONDS = 10.0 # avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle. _PERSIST_CHUNK_SIZE = 500 +# Admin mutations enqueue provider IDs here instead of running OpenRouter's +# per-model endpoint fan-out inside the request. One worker serializes refreshes +# and coalesces repeated mutations for the same provider. +_scheduled_provider_refresh_ids: set[int] = set() +_scheduled_provider_refresh_task: asyncio.Task[None] | None = None + # Visibility key used across this module: routing carries the provider # dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)), # so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps. @@ -86,9 +92,9 @@ class ProviderPathSnapshot: preserve_model_ids: frozenset[str] = frozenset() -def encode_model_path(provider_id: int, endpoint_tag: str | None = None) -> str: +def encode_model_path(provider_slug: str, endpoint_tag: str | None = None) -> str: """Encode a stable opaque selector without exposing upstream URLs.""" - components: list[tuple[str, str | int]] = [("provider", provider_id)] + components = [("provider", provider_slug)] if endpoint_tag: components.append(("endpoint", endpoint_tag)) return urlencode(components) @@ -370,7 +376,7 @@ async def _collect_provider_paths( def _base_path(model: object) -> DiscoveredPath: return DiscoveredPath( model_id=exposed_model_id(model), - path=encode_model_path(provider_identity.id), + path=encode_model_path(provider_identity.slug), provider=provider_identity, ) @@ -404,7 +410,7 @@ async def _collect_provider_paths( paths.extend( DiscoveredPath( model_id=model_id, - path=encode_model_path(provider_identity.id, endpoint.tag), + path=encode_model_path(provider_identity.slug, endpoint.tag), provider=provider_identity, endpoint_tag=endpoint.tag, endpoint_name=endpoint.provider_name, @@ -457,21 +463,31 @@ async def _persist_provider_paths( await session.exec(delete_stmt) # type: ignore[call-overload] for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE): chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE] + values = [ + { + "model_id": discovered.model_id, + "path": discovered.path, + "provider_slug": discovered.provider.slug, + "provider_type": discovered.provider.provider_type, + "endpoint_tag": discovered.endpoint_tag, + "endpoint_name": discovered.endpoint_name, + "upstream_provider_id": upstream_provider_id, + "updated_at": now, + } + for discovered in chunk + ] + insert_stmt = insert(ModelPathRow).values(values) await session.execute( - insert(ModelPathRow), - [ - { - "model_id": discovered.model_id, - "path": discovered.path, - "provider_slug": discovered.provider.slug, - "provider_type": discovered.provider.provider_type, - "endpoint_tag": discovered.endpoint_tag, - "endpoint_name": discovered.endpoint_name, - "upstream_provider_id": upstream_provider_id, - "updated_at": now, - } - for discovered in chunk - ], + insert_stmt.on_conflict_do_update( + index_elements=["model_id", "path", "upstream_provider_id"], + set_={ + "provider_slug": insert_stmt.excluded.provider_slug, + "provider_type": insert_stmt.excluded.provider_type, + "endpoint_tag": insert_stmt.excluded.endpoint_tag, + "endpoint_name": insert_stmt.excluded.endpoint_name, + "updated_at": insert_stmt.excluded.updated_at, + }, + ) ) await session.commit() @@ -560,7 +576,10 @@ async def refresh_model_paths( async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None: - """Immediately synchronize discovery after an admin provider/model mutation.""" + """Synchronize one provider when model-path discovery is enabled.""" + if _refresh_interval_seconds() <= 0: + return + from ..proxy import get_upstreams matching = [ @@ -574,6 +593,55 @@ async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None: await prune_model_paths_for_inactive_providers() +async def _drain_scheduled_provider_refreshes() -> None: + """Serialize and coalesce model-path refreshes scheduled by admin writes.""" + global _scheduled_provider_refresh_task + + try: + # Let mutations in the same event-loop turn collapse into one refresh. + await asyncio.sleep(0) + while _scheduled_provider_refresh_ids: + if _refresh_interval_seconds() <= 0: + _scheduled_provider_refresh_ids.clear() + return + provider_id = min(_scheduled_provider_refresh_ids) + _scheduled_provider_refresh_ids.remove(provider_id) + try: + await refresh_model_paths_for_provider(provider_id) + except asyncio.CancelledError: + raise + except Exception as exc: # noqa: BLE001 - background best effort + logger.warning( + "Failed to refresh model paths after admin mutation", + extra={ + "upstream_provider_id": provider_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) + finally: + _scheduled_provider_refresh_task = None + + +async def schedule_model_paths_refresh_for_provider( + upstream_provider_id: int, +) -> None: + """Queue a non-blocking, coalesced refresh after an admin mutation.""" + global _scheduled_provider_refresh_task + + if _refresh_interval_seconds() <= 0: + return + _scheduled_provider_refresh_ids.add(upstream_provider_id) + if ( + _scheduled_provider_refresh_task is None + or _scheduled_provider_refresh_task.done() + ): + _scheduled_provider_refresh_task = asyncio.create_task( + _drain_scheduled_provider_refreshes(), + name="model-path-admin-refresh", + ) + + def _refresh_interval_seconds() -> int: """Current interval, re-read every loop so runtime setting changes apply.""" from ..core.settings import settings diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index fe92c4f9..1caeaa5c 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -18,23 +18,6 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): supports_anthropic_messages = True litellm_provider_prefix = "openrouter/" - def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: - """Mirror ``_apply_provider_field``: strip repeated prefixes, map a - missing or self-echoing sub-provider to the literal ``"unknown"``.""" - provider_type = (self.provider_type or "").strip() - sub = (sub_provider or "").strip() - prefix = f"{provider_type}:" - while sub.lower().startswith(prefix.lower()): - sub = sub[len(prefix) :].strip() - if not sub or sub.lower() == provider_type.lower(): - return "unknown" - return f"{provider_type}:{sub}" - - def discovery_base_paths(self) -> list[str]: - """Native OpenRouter never stamps a bare ``openrouter``; a response - with no sub-provider is stamped ``unknown``.""" - return ["unknown"] - def _apply_provider_field(self, response_json: object) -> None: """Stamp the ``provider`` field for OpenRouter responses. diff --git a/tests/unit/test_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index c1ce13c8..17be72ec 100644 --- a/tests/unit/test_fee_payout_migration.py +++ b/tests/unit/test_fee_payout_migration.py @@ -4,6 +4,9 @@ import subprocess import sys from pathlib import Path +from alembic.config import Config +from alembic.script import ScriptDirectory + def _run_alembic(root: Path, database_url: str, revision: str) -> None: env = os.environ.copy() @@ -37,9 +40,10 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None: "payout_in_progress_msats, payout_started_at FROM routstr_fees" ).fetchone() - # Head of the 7f2843d3f4e4 lineage: model-paths chains onto the fee-payout - # repair migration. - assert version == ("4e0c3d195a49",) + migration_config = Config(str(root / "alembic.ini")) + assert version == ( + ScriptDirectory.from_config(migration_config).get_current_head(), + ) assert { "id", "accumulated_msats", diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 0538ea55..70986dfd 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -225,7 +225,7 @@ def _path_entry( if endpoint_tag or endpoint_name: endpoint = {"tag": endpoint_tag, "name": endpoint_name} return { - "path": mp.encode_model_path(provider_id, endpoint_tag), + "path": mp.encode_model_path(provider_slug or f"p{provider_id}", endpoint_tag), "provider": { "id": provider_id, "slug": provider_slug or f"p{provider_id}", @@ -253,10 +253,10 @@ def test_native_anthropic_not_openrouter() -> None: assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False -def test_encode_model_path_uses_provider_id_without_exposing_url() -> None: - assert mp.encode_model_path(42) == "provider=42" - assert mp.encode_model_path(42, "google-vertex/us-east5") == ( - "provider=42&endpoint=google-vertex%2Fus-east5" +def test_encode_model_path_uses_provider_slug_without_exposing_url() -> None: + assert mp.encode_model_path("openrouter-main") == "provider=openrouter-main" + assert mp.encode_model_path("openrouter-main", "google-vertex/us-east5") == ( + "provider=openrouter-main&endpoint=google-vertex%2Fus-east5" ) @@ -305,21 +305,6 @@ def test_openrouter_author_slug_none_when_no_slash() -> None: assert mp.openrouter_author_slug(m) is None -def test_discovery_paths_mirror_response_stamping() -> None: - """The discovery hook and ``_apply_provider_field`` must agree.""" - generic = _FakeProvider(provider_type="generic", base_url="https://x", models=[]) - assert generic.discovery_path_for_subprovider("Anthropic") == "generic:Anthropic" - assert generic.discovery_base_paths() == ["generic"] - - native = _FakeOpenRouterProvider(models=[]) - assert native.discovery_path_for_subprovider("GMICloud") == "openrouter:GMICloud" - # Sub-provider echoing the router name is stamped "unknown" on responses. - assert native.discovery_path_for_subprovider("OpenRouter") == "unknown" - assert native.discovery_path_for_subprovider("openrouter:openrouter") == "unknown" - assert native.discovery_path_for_subprovider(None) == "unknown" - assert native.discovery_base_paths() == ["unknown"] - - # --------------------------------------------------------------------------- # # Refresh through the public entry point # --------------------------------------------------------------------------- # @@ -412,7 +397,7 @@ async def test_disabling_model_on_one_provider_keeps_other_provider( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path("p1")} @pytest.mark.asyncio @@ -446,8 +431,8 @@ async def test_override_alias_not_applied_across_providers( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} - assert _paths_of(payload, "private-alias") == {mp.encode_model_path(2)} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path("p1")} + assert _paths_of(payload, "private-alias") == {mp.encode_model_path("p2")} @pytest.mark.asyncio @@ -699,9 +684,9 @@ async def test_openrouter_provider_adds_endpoint_paths( payload = await mp.get_paths_for_model("claude-opus-4.6") assert {item["path"] for item in payload["data"]} == { - mp.encode_model_path(2), - mp.encode_model_path(2, "google-vertex/eu"), - mp.encode_model_path(2, "google-vertex/us"), + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "google-vertex/eu"), + mp.encode_model_path("p2", "google-vertex/us"), } assert { item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] @@ -726,7 +711,10 @@ async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert paths == {mp.encode_model_path(2), mp.encode_model_path(2, "openrouter")} + assert paths == { + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "openrouter"), + } @pytest.mark.asyncio @@ -745,7 +733,10 @@ async def test_generic_provider_with_openrouter_base_url_discovers( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert paths == {mp.encode_model_path(1), mp.encode_model_path(1, "anthropic")} + assert paths == { + mp.encode_model_path("p1"), + mp.encode_model_path("p1", "anthropic"), + } @pytest.mark.asyncio @@ -776,6 +767,37 @@ async def test_openrouter_partial_failure_keeps_failed_models_previous_rows( assert _paths_of(await mp.get_all_model_paths(), "good") != before +@pytest.mark.asyncio +async def test_partial_failure_upserts_collapsed_public_model_paths( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A degraded canonical sibling may preserve the same public path that a + successful sibling refreshes; persistence must merge instead of rolling back.""" + provider = _FakeOpenRouterProvider( + models=[ + _model("vendora/shared", canonical_slug="vendora/shared"), + _model("vendorb/shared", canonical_slug="vendorb/shared"), + ], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + + def _one_sibling_degrades(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/vendorb/shared/endpoints"): + return httpx.Response(503) + return _endpoints_response("Google") + + _mock_transport(monkeypatch, _one_sibling_degrades) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "shared") == { + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "anthropic"), + mp.encode_model_path("p2", "google"), + } + + @pytest.mark.asyncio async def test_openrouter_failure_keeps_previous_rows( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch @@ -789,7 +811,7 @@ async def test_openrouter_failure_keeps_previous_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert mp.encode_model_path(2, "anthropic") in before + assert mp.encode_model_path("p2", "anthropic") in before def _network_down(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("network down", request=request) @@ -812,7 +834,10 @@ async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) await mp.refresh_model_paths([provider]) - expected = {mp.encode_model_path(2), mp.encode_model_path(2, "anthropic")} + expected = { + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "anthropic"), + } assert _paths_of(await mp.get_all_model_paths(), "m0") == expected counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) @@ -881,10 +906,10 @@ async def test_openrouter_shared_base_url_fetched_once( assert counter["requests"] == 1 paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path(2), - mp.encode_model_path(2, "anthropic"), - mp.encode_model_path(4), - mp.encode_model_path(4, "anthropic"), + mp.encode_model_path("p2"), + mp.encode_model_path("p2", "anthropic"), + mp.encode_model_path("p4"), + mp.encode_model_path("p4", "anthropic"), } @@ -947,8 +972,8 @@ async def test_same_model_two_providers_two_paths( entry = payload["data"][0] assert entry["id"] == "claude-opus-4.6" assert {p["path"] for p in entry["paths"]} == { - mp.encode_model_path(1), - mp.encode_model_path(2), + mp.encode_model_path("p1"), + mp.encode_model_path("p2"), } assert "canonical_id" not in entry assert all("canonical_id" not in p for p in entry["paths"]) @@ -1101,6 +1126,53 @@ async def test_refresh_model_paths_for_provider_selects_mutated_provider( assert seen == [[target]] +@pytest.mark.asyncio +async def test_admin_refresh_is_disabled_by_model_paths_kill_switch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + calls: list[int] = [] + + async def _fake_refresh(provider_id: int) -> None: + calls.append(provider_id) + + monkeypatch.setattr(mp, "refresh_model_paths_for_provider", _fake_refresh) + await mp.schedule_model_paths_refresh_for_provider(2) + await asyncio.sleep(0) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_admin_refresh_is_backgrounded_and_coalesced( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + mp._scheduled_provider_refresh_ids.clear() + mp._scheduled_provider_refresh_task = None + calls: list[int] = [] + + async def _fake_refresh(provider_id: int) -> None: + calls.append(provider_id) + + monkeypatch.setattr(mp, "refresh_model_paths_for_provider", _fake_refresh) + + await mp.schedule_model_paths_refresh_for_provider(2) + await mp.schedule_model_paths_refresh_for_provider(2) + task = mp._scheduled_provider_refresh_task + assert task is not None + await task + + assert calls == [2] + + @pytest.mark.asyncio async def test_refresh_loop_rereads_interval_and_picks_up_providers( monkeypatch: pytest.MonkeyPatch,