From 349d8dd009750798b50eda5675c9cb175e7fc673 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 5 Jul 2026 23:55:49 +0200 Subject: [PATCH 01/46] add model path endpoint --- docs/api/endpoints.md | 52 ++ docs/api/overview.md | 3 +- docs/provider/configuration.md | 6 + .../d7e8f9a0b1c2_add_model_paths_table.py | 66 ++ routstr/core/db.py | 35 ++ routstr/core/main.py | 11 + routstr/core/settings.py | 3 + routstr/payment/models.py | 23 + routstr/upstream/model_paths.py | 348 +++++++++++ tests/unit/test_model_paths.py | 591 ++++++++++++++++++ 10 files changed, 1137 insertions(+), 1 deletion(-) create mode 100644 migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py create mode 100644 routstr/upstream/model_paths.py create mode 100644 tests/unit/test_model_paths.py diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 5a9ea68d..efe57a68 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -327,6 +327,58 @@ GET /v1/models } ``` +### List Model Paths + +Get the upstream provider paths each advertised model can be reached through. +This is discovery data only; routing still chooses the provider per request. + +```http +GET /v1/models/paths +``` + +**Response:** + +```json +{ + "data": [ + { + "id": "claude-sonnet-4", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"} + ] + } + ] +} +``` + +### List Paths for One Model + +Use a query parameter so model IDs containing `/` are handled safely. Lookup is +by the public, unqualified model ID: `glm-5v-turbo` resolves +`z-ai/glm-5v-turbo`, and `deepseek-v4-pro` and `deepseek/deepseek-v4-pro` +return the same merged path set. + +```http +GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 +``` + +**Response:** + +```json +{ + "data": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"} + ] +} +``` + +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`. + ## Wallet Management ### Create Wallet (Coming Soon) diff --git a/docs/api/overview.md b/docs/api/overview.md index 92fedd7e..c82e4e22 100644 --- a/docs/api/overview.md +++ b/docs/api/overview.md @@ -100,6 +100,7 @@ All errors follow a consistent format: Standard OpenAI-compatible endpoints: - **Models**: `/v1/models` +- **Model paths**: `/v1/models/paths`, `/v1/models/paths/model?model_id=...` - **Responses**: `/v1/responses` - **Chat Completions**: `/v1/chat/completions` - **Embeddings**: `/v1/embeddings` @@ -302,7 +303,7 @@ Get node metadata: GET /v1/info ``` -Supported models and pricing are available at `/v1/models`. +Supported models and pricing are available at `/v1/models`. Upstream provider path discovery is available at `/v1/models/paths` and `/v1/models/paths/model?model_id=...`. ## Next Steps diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 40930676..45387fb4 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -137,6 +137,7 @@ 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` | ### Priority @@ -156,3 +157,8 @@ Manage which AI models you offer: - **Create aliases** — friendly names for models See [Pricing](pricing.md) for per-model pricing strategies. + +Model path discovery is refreshed in the background and exposed through +`/v1/models/paths`. The response groups each client-visible model ID with the +provider paths that may appear in chat-completion response metadata. Tune the +refresh cadence with `MODEL_PATHS_REFRESH_INTERVAL_SECONDS`. diff --git a/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py b/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py new file mode 100644 index 00000000..e5ea85f4 --- /dev/null +++ b/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py @@ -0,0 +1,66 @@ +"""add model_paths table + +Revision ID: d7e8f9a0b1c2 +Revises: c6d7e8f9a0b1 +Create Date: 2026-07-05 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "d7e8f9a0b1c2" +down_revision = "c6d7e8f9a0b1" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + if "model_paths" in inspector.get_table_names(): + return + + op.create_table( + "model_paths", + sa.Column("id", sa.Integer(), primary_key=True, autoincrement=True), + sa.Column("model_id", sa.String(), nullable=False), + sa.Column("path", sa.String(), nullable=False), + sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint( + ["upstream_provider_id"], + ["upstream_providers.id"], + ondelete="CASCADE", + ), + sa.UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + op.create_index( + "ix_model_paths_model_id", + "model_paths", + ["model_id"], + ) + op.create_index( + "ix_model_paths_upstream_provider_id", + "model_paths", + ["upstream_provider_id"], + ) + + +def downgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + if "model_paths" not in inspector.get_table_names(): + return + + existing_indexes = {idx["name"] for idx in inspector.get_indexes("model_paths")} + if "ix_model_paths_upstream_provider_id" in existing_indexes: + op.drop_index("ix_model_paths_upstream_provider_id", table_name="model_paths") + if "ix_model_paths_model_id" in existing_indexes: + op.drop_index("ix_model_paths_model_id", table_name="model_paths") + op.drop_table("model_paths") diff --git a/routstr/core/db.py b/routstr/core/db.py index 586f467e..4023c0c9 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -210,6 +210,41 @@ class ModelRow(SQLModel, table=True): # type: ignore upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") +class ModelPathRow(SQLModel, table=True): # type: ignore + """Upstream provider path a model is reachable through. + + Discovery/visibility data only. ``model_id`` is intentionally NOT globally + unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or + id``) grouped across every provider that exposes the model. A single model + can therefore have several rows — one per direct provider path plus one per + OpenRouter sub-provider endpoint. + """ + + __tablename__ = "model_paths" + __table_args__ = ( + UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + id: int | None = Field(default=None, primary_key=True) + model_id: str = Field( + index=True, description="Client-visible /v1/models id (forwarded_model_id or id)" + ) + path: str = Field( + description="Provider path stamped on chat completion responses, e.g. " + "'anthropic' or 'openrouter:Anthropic'" + ) + upstream_provider_id: int = Field( + index=True, + foreign_key="upstream_providers.id", + ondelete="CASCADE", + description="upstream_providers.id this path was discovered from", + ) + + class LightningInvoice(SQLModel, table=True): # type: ignore __tablename__ = "lightning_invoices" diff --git a/routstr/core/main.py b/routstr/core/main.py index 903d6d1a..d105331a 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -58,6 +58,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: providers_task = None models_refresh_task = None model_maps_refresh_task = None + model_paths_refresh_task = None key_reset_task = None stale_reservation_task = None dead_key_prune_task = None @@ -127,6 +128,12 @@ 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 + + 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()) @@ -173,6 +180,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: models_refresh_task.cancel() if model_maps_refresh_task is not None: model_maps_refresh_task.cancel() + if model_paths_refresh_task is not None: + model_paths_refresh_task.cancel() if key_reset_task is not None: key_reset_task.cancel() if stale_reservation_task is not None: @@ -206,6 +215,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(models_refresh_task) if model_maps_refresh_task is not None: tasks_to_wait.append(model_maps_refresh_task) + if model_paths_refresh_task is not None: + tasks_to_wait.append(model_paths_refresh_task) if key_reset_task is not None: tasks_to_wait.append(key_reset_task) if stale_reservation_task is not None: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 3a144e10..5fce93af 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -94,6 +94,9 @@ class Settings(BaseSettings): models_refresh_interval_seconds: int = Field( default=360, env="MODELS_REFRESH_INTERVAL_SECONDS" ) + model_paths_refresh_interval_seconds: int = Field( + default=600, env="MODEL_PATHS_REFRESH_INTERVAL_SECONDS" + ) enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") diff --git a/routstr/payment/models.py b/routstr/payment/models.py index e6ea643d..b21669a8 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -596,6 +596,29 @@ async def test_model( } +@models_router.get("/v1/models/paths") +@models_router.get("/v1/models/paths/", include_in_schema=False) +async def model_paths() -> dict: + """All models with every upstream provider path they are reachable through.""" + from ..upstream.model_paths import get_all_model_paths + + return {"data": await get_all_model_paths()} + + +@models_router.get("/v1/models/paths/model") +@models_router.get("/v1/models/paths/model/", include_in_schema=False) +async def model_paths_for_model(model_id: str) -> dict: + """Paths for a single model. + + Uses a query parameter (``?model_id=...``) under a fully static route so + model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL + encoding and there is no dynamic-route ambiguity. + """ + from ..upstream.model_paths import get_paths_for_model + + return {"data": await get_paths_for_model(model_id)} + + @models_router.get("/v1/models") @models_router.get("/v1/models/", include_in_schema=False) @models_router.get("/models") diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py new file mode 100644 index 00000000..5ce1f527 --- /dev/null +++ b/routstr/upstream/model_paths.py @@ -0,0 +1,348 @@ +"""Model-path discovery service. + +Exposes every upstream provider path a Routstr model is reachable through. +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``): + +- Direct upstream -> ```` e.g. ``anthropic`` +- Generic/custom OpenRouter-compatible upstream -> ``generic:`` +- Native OpenRouter routing to a sub-provider -> ``openrouter:`` + +Native OpenRouter does not emit a useful bare ``openrouter`` path when no +sub-provider is present; it reports ``unknown`` instead. +""" + +from __future__ import annotations + +import asyncio +import random +from typing import Callable + +import httpx +from sqlmodel import col, delete, select + +from ..core.db import ModelPathRow, create_session +from ..core.logging import get_logger +from .base import BaseUpstreamProvider + +logger = get_logger(__name__) + +# Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds +# of models does not open hundreds of concurrent requests every refresh. +_OPENROUTER_CONCURRENCY = 5 +_OPENROUTER_TIMEOUT_SECONDS = 10.0 + + +def is_openrouter_base_url(base_url: str | None) -> bool: + """True when ``base_url`` points at OpenRouter. + + Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``: + that predicate also returns True for native Anthropic (correct for + cache-control, wrong for OpenRouter endpoint discovery). This one keys only + on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched + while native Anthropic is not. + """ + return "openrouter.ai" in (base_url or "") + + +def exposed_model_id(model: object) -> str: + """Client-visible ``/v1/models`` id for a cached model.""" + forwarded = getattr(model, "forwarded_model_id", None) + return forwarded or getattr(model, "id") + + +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``. + """ + return model_id.rsplit("/", 1)[-1] + + +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. + """ + canonical = getattr(model, "canonical_slug", None) + if canonical and "/" in canonical: + return canonical + model_id = getattr(model, "id", None) + if model_id and "/" in model_id: + return model_id + return None + + +async def _fetch_openrouter_endpoint_paths( + client: httpx.AsyncClient, + base_url: str, + api_key: str, + author_slug: str, + path_prefix: str, + semaphore: asyncio.Semaphore, +) -> list[str]: + """Return ``:`` paths for one model, or ``[]``. + + Failures (network, rate limit, bad payload) are logged and swallowed so one + model never breaks the whole refresh. + """ + url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" + headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} + async with semaphore: + try: + resp = await client.get( + url, headers=headers, timeout=_OPENROUTER_TIMEOUT_SECONDS + ) + except Exception as e: # noqa: BLE001 - isolate per-model failures + logger.warning( + "OpenRouter endpoint discovery request failed", + extra={"author_slug": author_slug, "error": str(e)}, + ) + return [] + + if resp.status_code == 429: + logger.warning( + "OpenRouter endpoint discovery rate-limited", + extra={"author_slug": author_slug}, + ) + return [] + if resp.status_code != 200: + logger.warning( + "OpenRouter endpoint discovery non-200", + extra={"author_slug": author_slug, "status_code": resp.status_code}, + ) + return [] + + try: + endpoints = resp.json().get("data", {}).get("endpoints", []) + except Exception as e: # noqa: BLE001 + logger.warning( + "OpenRouter endpoint discovery bad payload", + extra={"author_slug": author_slug, "error": str(e)}, + ) + return [] + + 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)) + + +async def _collect_provider_paths( + upstream: BaseUpstreamProvider, +) -> list[tuple[str, str]]: + """Collect ``(model_id, path)`` pairs for one provider instance. + + Emits the direct ```` path for normal upstreams. For + OpenRouter-compatible providers, emits one path per OpenRouter sub-provider + endpoint, prefixed the same way response stamping prefixes it. + """ + provider_type = (upstream.provider_type or "").strip() + models = [m for m in upstream.get_cached_models() if getattr(m, "enabled", True)] + is_openrouter = is_openrouter_base_url(upstream.base_url) + + 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 provider_type: + return pairs + + semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) + async with httpx.AsyncClient() as client: + + async def _for_model(model: object) -> list[tuple[str, str]]: + author_slug = openrouter_author_slug(model) + if not author_slug: + return [] + paths = await _fetch_openrouter_endpoint_paths( + client, + upstream.base_url, + upstream.api_key, + author_slug, + provider_type, + semaphore, + ) + model_id = exposed_model_id(model) + return [(model_id, path) for path in paths] + + results = await asyncio.gather( + *(_for_model(m) for m in models), return_exceptions=True + ) + + for result in results: + if isinstance(result, BaseException): + logger.warning( + "OpenRouter endpoint discovery task errored", + extra={"provider": provider_type, "error": str(result)}, + ) + continue + pairs.extend(result) + + return pairs + + +async def _persist_provider_paths( + upstream_provider_id: int, pairs: list[tuple[str, str]] +) -> None: + """Replace all rows for ``upstream_provider_id`` with ``pairs``. + + Replacement (not upsert) so stale paths disappear when provider config or + upstream availability changes. + """ + unique_pairs = list(dict.fromkeys(pairs)) + 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, + ) + ) + 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 with create_session() as session: + stmt = delete(ModelPathRow) + if active_provider_ids: + stmt = stmt.where( + col(ModelPathRow.upstream_provider_id).not_in(active_provider_ids) + ) + await session.exec(stmt) # type: ignore[call-overload] + await session.commit() + + +async def refresh_model_paths( + upstreams: list[BaseUpstreamProvider], +) -> None: + """Recompute and persist model paths for every enabled provider. + + One provider's failure is logged and isolated; it must not break the rest. + """ + active_provider_ids = { + upstream.db_id for upstream in upstreams if upstream.db_id is not None + } + await _prune_inactive_provider_paths(active_provider_ids) + + for upstream in upstreams: + if upstream.db_id is None: + continue + try: + pairs = await _collect_provider_paths(upstream) + await _persist_provider_paths(upstream.db_id, pairs) + except Exception as e: # noqa: BLE001 - isolate per-provider failures + logger.error( + "Failed to refresh model paths for provider", + extra={ + "provider": upstream.provider_type or upstream.base_url, + "db_id": upstream.db_id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + + +async def refresh_model_paths_periodically( + upstreams_provider: ( + Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider] + ), +) -> None: + """Background task mirroring ``refresh_upstreams_models_periodically``.""" + from ..core.settings import settings + + interval = getattr(settings, "model_paths_refresh_interval_seconds", 0) + if not interval or interval <= 0: + logger.info("Model paths refresh disabled (interval <= 0)") + return + + def _resolve_upstreams() -> list[BaseUpstreamProvider]: + if callable(upstreams_provider): + return upstreams_provider() + return upstreams_provider + + while True: + try: + await refresh_model_paths(_resolve_upstreams()) + except asyncio.CancelledError: + break + except Exception as e: # noqa: BLE001 + logger.error( + "Error in model paths refresh loop", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + try: + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) + except asyncio.CancelledError: + break + + +async def get_all_model_paths() -> list[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)) + ).all() + + grouped: dict[str, list[dict]] = {} + seen_paths: dict[str, set[str]] = {} + for row in rows: + 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()] + + +async def get_paths_for_model(model_id: str) -> list[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. + """ + requested_id = public_model_id(model_id) + async with create_session() as session: + rows = ( + await session.exec( + select(ModelPathRow).order_by( + ModelPathRow.path, + col(ModelPathRow.upstream_provider_id), + ModelPathRow.model_id, + ) + ) + ).all() + + seen: set[str] = set() + paths: list[dict] = [] + for row in rows: + if public_model_id(row.model_id) != requested_id: + continue + if row.path in seen: + continue + seen.add(row.path) + paths.append({"path": row.path}) + return paths diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py new file mode 100644 index 00000000..75468383 --- /dev/null +++ b/tests/unit/test_model_paths.py @@ -0,0 +1,591 @@ +"""Tests for the model-path discovery service and endpoints.""" + +from __future__ import annotations + +import asyncio +import os +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Any, AsyncGenerator + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.db import UpstreamProviderRow # noqa: E402 +from routstr.payment.models import models_router # noqa: E402 +from routstr.upstream import model_paths as mp # noqa: E402 + +# --------------------------------------------------------------------------- # +# Fakes +# --------------------------------------------------------------------------- # + + +def _model( + id: str, + *, + forwarded_model_id: str | None = None, + canonical_slug: str | None = None, + enabled: bool = True, +) -> SimpleNamespace: + return SimpleNamespace( + id=id, + forwarded_model_id=forwarded_model_id, + canonical_slug=canonical_slug, + enabled=enabled, + ) + + +class _FakeProvider: + def __init__( + self, + *, + provider_type: str, + base_url: str, + models: list[SimpleNamespace], + db_id: int | None = 1, + api_key: str = "sk-test", + ) -> None: + self.provider_type = provider_type + self.base_url = base_url + self.api_key = api_key + self.db_id = db_id + self._models = models + + def get_cached_models(self) -> list[SimpleNamespace]: + return self._models + + +class _FakeResponse: + def __init__(self, status_code: int, payload: Any) -> None: + self.status_code = status_code + self._payload = payload + + def json(self) -> Any: + return self._payload + + +@pytest.fixture +async def patched_session( + monkeypatch: pytest.MonkeyPatch, +) -> AsyncGenerator[AsyncEngine, None]: + """Bind the service's ``create_session`` to a fresh in-memory engine.""" + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + # Seed the FK target so ModelPathRow inserts satisfy the constraint. + async with AsyncSession(engine) as session: + for pid in (1, 2): + session.add( + UpstreamProviderRow( + id=pid, + slug=f"p{pid}", + provider_type="anthropic" if pid == 1 else "openrouter", + base_url=f"https://provider-{pid}", + api_key=f"k{pid}", + ) + ) + await session.commit() + + @asynccontextmanager + async def _factory() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + monkeypatch.setattr(mp, "create_session", _factory) + yield engine + await engine.dispose() + + +# --------------------------------------------------------------------------- # +# Predicates / pure helpers +# --------------------------------------------------------------------------- # + + +def test_is_openrouter_base_url() -> None: + assert mp.is_openrouter_base_url("https://openrouter.ai/api/v1") is True + assert mp.is_openrouter_base_url("https://api.anthropic.com") is False + assert mp.is_openrouter_base_url(None) is False + + +def test_native_anthropic_not_openrouter() -> None: + """Native Anthropic must not be treated as OpenRouter-compatible even though + ``_upstream_accepts_cache_control`` returns True for it.""" + assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False + + +def test_exposed_model_id_prefers_forwarded() -> None: + assert ( + mp.exposed_model_id(_model("claude-x", forwarded_model_id="fwd-claude")) + == "fwd-claude" + ) + assert mp.exposed_model_id(_model("claude-x")) == "claude-x" + + +def test_public_model_id_strips_provider_prefix() -> None: + assert mp.public_model_id("z-ai/glm-5v-turbo") == "glm-5v-turbo" + assert mp.public_model_id("gpt-4o-mini") == "gpt-4o-mini" + + +def test_openrouter_author_slug_uses_canonical_not_forwarded() -> None: + m = _model( + "claude-opus-4.6", + forwarded_model_id="forwarded-only", + canonical_slug="anthropic/claude-opus-4.6", + ) + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_falls_back_to_slash_id() -> None: + m = _model("anthropic/claude-opus-4.6", canonical_slug="claude-opus-4.6") + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_none_when_no_slash() -> None: + m = _model("claude-opus-4.6", canonical_slug="claude-opus-4.6") + assert mp.openrouter_author_slug(m) is None + + +# --------------------------------------------------------------------------- # +# Collection +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_direct_provider_single_path_uses_provider_type() -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + ) + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("claude-opus-4.6", "anthropic")] + + +@pytest.mark.asyncio +async def test_direct_path_stores_exposed_model_id() -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("internal-id", forwarded_model_id="claude-opus-4.6")], + ) + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("claude-opus-4.6", "anthropic")] + + +@pytest.mark.asyncio +async def test_disabled_models_excluded() -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[ + _model("enabled-model"), + _model("disabled-model", enabled=False), + ], + ) + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("enabled-model", "anthropic")] + + +@pytest.mark.asyncio +async def test_openrouter_provider_adds_endpoint_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _fake_get( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + return _FakeResponse( + 200, + { + "data": { + "endpoints": [ + {"provider_name": "Anthropic"}, + {"provider_name": "Amazon Bedrock"}, + ] + } + }, + ) + + monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert ("claude-opus-4.6", "openrouter:Anthropic") in pairs + assert ("claude-opus-4.6", "openrouter:Amazon Bedrock") in pairs + assert ("claude-opus-4.6", "openrouter") not in pairs + assert len(pairs) == 2 + + +@pytest.mark.asyncio +async def test_generic_provider_with_openrouter_base_url_discovers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A generic provider pointed at OpenRouter exposes the response-stamped + ``generic:`` path, not a native ``openrouter:`` path.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _fake_get( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + return _FakeResponse( + 200, {"data": {"endpoints": [{"provider_name": "Anthropic"}]}} + ) + + monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [("claude-opus-4.6", "generic:Anthropic")] + + +@pytest.mark.asyncio +async def test_openrouter_failure_degrades_gracefully( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _boom( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + raise RuntimeError("network down") + + monkeypatch.setattr("httpx.AsyncClient.get", _boom) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [] + + +@pytest.mark.asyncio +async def test_openrouter_rate_limit_skips_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + ) + + async def _rate_limited( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + return _FakeResponse(429, {}) + + monkeypatch.setattr("httpx.AsyncClient.get", _rate_limited) + + pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert pairs == [] + + +@pytest.mark.asyncio +async def test_openrouter_fanout_is_bounded( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) + models = [ + _model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20) + ] + provider = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=models, + ) + + state = {"current": 0, "max": 0} + + async def _slow_get( + self: object, + url: str, + headers: object = None, + timeout: object = None, + ) -> _FakeResponse: + state["current"] += 1 + state["max"] = max(state["max"], state["current"]) + await asyncio.sleep(0.02) + state["current"] -= 1 + return _FakeResponse(200, {"data": {"endpoints": [{"provider_name": "X"}]}}) + + monkeypatch.setattr("httpx.AsyncClient.get", _slow_get) + + await mp._collect_provider_paths(provider) # type: ignore[arg-type] + assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" + + +# --------------------------------------------------------------------------- # +# Persistence + query +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_refresh_replaces_stale_rows( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("m1", "anthropic"), ("m2", "anthropic")]) + first = await mp.get_all_model_paths() + assert {row["id"] for row in first} == {"m1", "m2"} + + # Second refresh with a different set — stale m2 must disappear. + await mp._persist_provider_paths(1, [("m1", "anthropic")]) + second = await mp.get_all_model_paths() + assert {row["id"] for row in second} == {"m1"} + + +@pytest.mark.asyncio +async def test_refresh_model_paths_prunes_inactive_provider_rows( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("m1", "anthropic")]) + await mp._persist_provider_paths(2, [("m2", "openrouter:Anthropic")]) + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + active_only = await mp.get_all_model_paths() + assert active_only == [{"id": "m1", "paths": [{"path": "anthropic"}]}] + + await mp.refresh_model_paths([]) + assert await mp.get_all_model_paths() == [] + + +@pytest.mark.asyncio +async def test_same_model_two_providers_two_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) + + data = await mp.get_all_model_paths() + assert len(data) == 1 + entry = data[0] + assert entry["id"] == "claude-opus-4.6" + paths = {p["path"] for p in entry["paths"]} + assert paths == {"anthropic", "openrouter:Anthropic"} + # No canonical_id anywhere. + assert "canonical_id" not in entry + assert all("canonical_id" not in p for p in entry["paths"]) + + +@pytest.mark.asyncio +async def test_get_all_model_paths_deduplicates_visible_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) + + assert await mp.get_all_model_paths() == [ + {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + ] + + +@pytest.mark.asyncio +async def test_get_all_model_paths_returns_unqualified_model_ids( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) + await mp._persist_provider_paths(5, [("openai/gpt-4o-mini", "openrouter:OpenAI")]) + + data = await mp.get_all_model_paths() + + assert {row["id"] for row in data} == {"glm-5v-turbo", "gpt-4o-mini"} + + +@pytest.mark.asyncio +async def test_get_paths_for_model_returns_only_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) + + paths = await mp.get_paths_for_model("claude-opus-4.6") + assert {p["path"] for p in paths} == {"anthropic", "openrouter:Anthropic"} + assert all(set(p.keys()) == {"path"} for p in paths) + assert await mp.get_paths_for_model("does-not-exist") == [] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) + + paths = await mp.get_paths_for_model("glm-5v-turbo") + + assert paths == [{"path": "openrouter:Z.AI"}] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_deduplicates_visible_paths( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) + await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) + + assert await mp.get_paths_for_model("claude-opus-4.6") == [ + {"path": "anthropic"} + ] + assert await mp.get_paths_for_model("anthropic/claude-opus-4.6") == [ + {"path": "anthropic"} + ] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(7, [("deepseek-v4-pro", "generic")]) + await mp._persist_provider_paths( + 4, [("deepseek/deepseek-v4-pro", "openrouter:DeepSeek")] + ) + + short_paths = await mp.get_paths_for_model("deepseek-v4-pro") + prefixed_paths = await mp.get_paths_for_model("deepseek/deepseek-v4-pro") + + assert short_paths == [ + {"path": "generic"}, + {"path": "openrouter:DeepSeek"}, + ] + assert prefixed_paths == short_paths + + +@pytest.mark.asyncio +async def test_refresh_model_paths_skips_provider_without_db_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=None, + ) + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + assert await mp.get_all_model_paths() == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_isolates_provider_failure( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + good = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + bad = _FakeProvider( + provider_type="openrouter", + base_url="https://openrouter.ai/api/v1", + models=[_model("m", canonical_slug="a/m")], + db_id=2, + ) + + original = mp._collect_provider_paths + + async def _maybe_fail(upstream: Any) -> list[tuple[str, str]]: + if upstream is bad: + raise RuntimeError("boom") + return await original(upstream) # type: ignore[arg-type] + + monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) + + await mp.refresh_model_paths([good, bad]) # type: ignore[list-item] + data = await mp.get_all_model_paths() + assert {row["id"] for row in data} == {"claude-opus-4.6"} + + +# --------------------------------------------------------------------------- # +# Endpoints +# --------------------------------------------------------------------------- # + + +def _make_model_paths_app() -> FastAPI: + app = FastAPI() + app.include_router(models_router) + return app + + +def test_model_paths_endpoint_returns_all_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def _fake_get_all_model_paths() -> list[dict[str, Any]]: + return [ + { + "id": "claude-opus-4.6", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"}, + ], + } + ] + + monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) + + response = TestClient(_make_model_paths_app()).get("/v1/models/paths") + + assert response.status_code == 200 + assert response.json() == { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"}, + ], + } + ] + } + + +def test_model_paths_for_model_endpoint_accepts_slash_model_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[str] = [] + + async def _fake_get_paths_for_model(model_id: str) -> list[dict[str, Any]]: + calls.append(model_id) + return [{"path": "generic:Anthropic"}] + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", + params={"model_id": "anthropic/claude-opus-4.6"}, + ) + + assert response.status_code == 200 + assert response.json() == {"data": [{"path": "generic:Anthropic"}]} + assert calls == ["anthropic/claude-opus-4.6"] From dc25659cff29aad9316344c6c5cf848fe6e95181 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 7 Jul 2026 11:11:46 +0200 Subject: [PATCH 02/46] only activ model should be visible --- routstr/upstream/model_paths.py | 136 ++++++++++++++++++++++++++++++-- tests/unit/test_model_paths.py | 127 ++++++++++++++++++++++++++++- 2 files changed, 254 insertions(+), 9 deletions(-) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 5ce1f527..94d21833 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -19,15 +19,19 @@ from __future__ import annotations import asyncio import random -from typing import Callable +from typing import TYPE_CHECKING, Callable import httpx +from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select -from ..core.db import ModelPathRow, create_session +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 + logger = get_logger(__name__) # Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds @@ -138,8 +142,117 @@ async def _fetch_openrouter_endpoint_paths( return list(dict.fromkeys(paths)) +async def _load_model_visibility() -> tuple[ + dict[str, tuple[ModelRow, float]], set[str], 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``. + """ + async with create_session() as session: + query = select(UpstreamProviderRow).options( + selectinload(UpstreamProviderRow.models) # type: ignore[arg-type] + ) + provider_rows = (await session.exec(query)).all() + + overrides_by_id: dict[str, tuple[ModelRow, float]] = {} + disabled_model_ids: set[str] = set() + enabled_provider_ids: set[int] = set() + + for provider in provider_rows: + if not provider.enabled: + continue + if provider.id is not None: + enabled_provider_ids.add(provider.id) + for model in provider.models: + if model.enabled: + overrides_by_id[model.id] = (model, provider.provider_fee) + else: + disabled_model_ids.add(model.id) + + 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 + + +def _apply_model_visibility( + upstream: BaseUpstreamProvider, + overrides_by_id: dict[str, tuple[ModelRow, float]] | None, + disabled_model_ids: set[str] | 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() + 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: + 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) + 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()) + + 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]]: """Collect ``(model_id, path)`` pairs for one provider instance. @@ -148,7 +261,7 @@ async def _collect_provider_paths( endpoint, prefixed the same way response stamping prefixes it. """ provider_type = (upstream.provider_type or "").strip() - models = [m for m in upstream.get_cached_models() if getattr(m, "enabled", True)] + models = _apply_model_visibility(upstream, overrides_by_id, disabled_model_ids) is_openrouter = is_openrouter_base_url(upstream.base_url) pairs: list[tuple[str, str]] = [] @@ -240,16 +353,27 @@ async def refresh_model_paths( One provider's failure is logged and isolated; it must not break the rest. """ + ( + overrides_by_id, + disabled_model_ids, + enabled_provider_ids, + ) = await _load_model_visibility() active_provider_ids = { - upstream.db_id for upstream in upstreams if upstream.db_id is not None + 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) for upstream in upstreams: - if upstream.db_id is None: + if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: continue try: - pairs = await _collect_provider_paths(upstream) + pairs = await _collect_provider_paths( + upstream, + overrides_by_id=overrides_by_id, + disabled_model_ids=disabled_model_ids, + ) await _persist_provider_paths(upstream.db_id, pairs) except Exception as e: # noqa: BLE001 - isolate per-provider failures logger.error( diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 75468383..947c4312 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import json import os from contextlib import asynccontextmanager from types import SimpleNamespace @@ -18,7 +19,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") os.environ.setdefault("UPSTREAM_API_KEY", "test") -from routstr.core.db import UpstreamProviderRow # noqa: E402 +from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402 from routstr.payment.models import models_router # noqa: E402 from routstr.upstream import model_paths as mp # noqa: E402 @@ -42,6 +43,37 @@ def _model( ) +def _model_row( + id: str, + *, + upstream_provider_id: int = 1, + forwarded_model_id: str | None = None, + canonical_slug: str | None = None, + enabled: bool = True, +) -> ModelRow: + return ModelRow( + id=id, + upstream_provider_id=upstream_provider_id, + name=id, + created=0, + description="test model", + context_length=8192, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "test", + "instruct_type": None, + } + ), + pricing=json.dumps({"prompt": 0.000001, "completion": 0.000002}), + enabled=enabled, + forwarded_model_id=forwarded_model_id, + canonical_slug=canonical_slug, + ) + + class _FakeProvider: def __init__( self, @@ -360,6 +392,70 @@ async def test_refresh_replaces_stale_rows( assert {row["id"] for row in second} == {"m1"} +@pytest.mark.asyncio +async def test_refresh_model_paths_excludes_db_disabled_override( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("disabled-by-db", enabled=False)) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("disabled-by-db")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_uses_db_forwarded_alias( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("internal-id", forwarded_model_id="public-alias")) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("internal-id")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [ + {"id": "public-alias", "paths": [{"path": "anthropic"}]} + ] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cache( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("deployment-id", forwarded_model_id="public-deployment")) + await session.commit() + + provider = _FakeProvider( + provider_type="generic", + base_url="https://custom-provider/v1", + models=[], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [ + {"id": "public-deployment", "paths": [{"path": "generic"}]} + ] + + @pytest.mark.asyncio async def test_refresh_model_paths_prunes_inactive_provider_rows( patched_session: AsyncEngine, @@ -382,6 +478,29 @@ async def test_refresh_model_paths_prunes_inactive_provider_rows( assert await mp.get_all_model_paths() == [] +@pytest.mark.asyncio +async def test_refresh_model_paths_skips_disabled_db_provider( + patched_session: AsyncEngine, +) -> None: + await mp._persist_provider_paths(1, [("stale-model", "anthropic")]) + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.enabled = False + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("fresh-model")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) # type: ignore[list-item] + + assert await mp.get_all_model_paths() == [] + + @pytest.mark.asyncio async def test_same_model_two_providers_two_paths( patched_session: AsyncEngine, @@ -515,10 +634,12 @@ async def test_refresh_model_paths_isolates_provider_failure( original = mp._collect_provider_paths - async def _maybe_fail(upstream: Any) -> list[tuple[str, str]]: + async def _maybe_fail( + upstream: Any, *args: Any, **kwargs: Any + ) -> list[tuple[str, str]]: if upstream is bad: raise RuntimeError("boom") - return await original(upstream) # type: ignore[arg-type] + return await original(upstream, *args, **kwargs) # type: ignore[arg-type] monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) From 7ed18a9d02733aa855dabc38cbd84b387d92b23d Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 10 Jul 2026 19:45:55 +0200 Subject: [PATCH 03/46] fix: per-mint rate limiting, trusted-mint fallback, and retry factory fix --- .../add_mint_url_to_lightning_invoices.py | 21 ++ routstr/core/db.py | 3 + routstr/core/settings.py | 17 + routstr/lightning.py | 84 ++++- routstr/payment/lnurl.py | 24 +- routstr/wallet.py | 329 ++++++++++++++++-- .../test_lightning_invoice_rip08.py | 3 +- tests/unit/test_wallet.py | 241 +++++++++++++ 8 files changed, 669 insertions(+), 53 deletions(-) create mode 100644 migrations/versions/add_mint_url_to_lightning_invoices.py diff --git a/migrations/versions/add_mint_url_to_lightning_invoices.py b/migrations/versions/add_mint_url_to_lightning_invoices.py new file mode 100644 index 00000000..d3eb71b9 --- /dev/null +++ b/migrations/versions/add_mint_url_to_lightning_invoices.py @@ -0,0 +1,21 @@ +"""add mint_url to lightning_invoices + +Revision ID: add_mint_url_li +Revises: c6d7e8f9a0b1 +Create Date: 2026-07-10 02:00:00.000000 +""" +import sqlalchemy as sa +from alembic import op + +revision = "add_mint_url_li" +down_revision = "c6d7e8f9a0b1" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column("lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True)) + + +def downgrade() -> None: + op.drop_column("lightning_invoices", "mint_url") diff --git a/routstr/core/db.py b/routstr/core/db.py index 586f467e..6cd5b593 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -225,6 +225,9 @@ class LightningInvoice(SQLModel, table=True): # type: ignore default=None, description="Associated API key hash for topup operations" ) purpose: str = Field(description="create or topup") + mint_url: str | None = Field( + default=None, description="Mint URL where the quote was created (fallback tracking)" + ) created_at: int = Field( default_factory=lambda: int(time.time()), description="Unix timestamp" ) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 3a144e10..88797ff2 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -49,6 +49,23 @@ class Settings(BaseSettings): payout_interval_seconds: int = Field( default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS" ) + # Timeout (seconds) for individual mint API operations (melt, mint, swap, + # checkstate). When a mint is slow or rate-limiting, operations are + # cancelled after this delay instead of hanging indefinitely. + mint_operation_timeout_seconds: int = Field( + default=30, gt=0, env="MINT_OPERATION_TIMEOUT_SECONDS" + ) + # Maximum mint API requests per minute, per mint URL. Nutshell mints + # (e.g. Minibits) enforce 20/min/IP on transaction endpoints (mint, melt, + # swap, quotes) and 60/min/IP globally. 20 stays under the transaction + # bucket since most calls here are transaction ops. 0 = unlimited. + mint_max_requests_per_minute: int = Field( + default=20, ge=0, env="MINT_MAX_REQUESTS_PER_MINUTE" + ) + # Max retries when a mint returns 429 or times out (exponential backoff). + mint_retry_max_attempts: int = Field( + default=3, ge=0, env="MINT_RETRY_MAX_ATTEMPTS" + ) # Pricing # Default behavior: derive pricing from MODELS diff --git a/routstr/lightning.py b/routstr/lightning.py index b0bbc63b..24c21dde 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -11,7 +11,13 @@ from sqlmodel.ext.asyncio.session import AsyncSession from .core.db import ApiKey, LightningInvoice, create_session, get_session from .core.logging import get_logger from .core.settings import settings -from .wallet import get_wallet +from .wallet import ( + MintConnectionError, + _is_mint_rate_limited, + _mint_operation, + get_wallet, + is_mint_connection_error, +) logger = get_logger(__name__) @@ -64,12 +70,44 @@ class InvoiceRecoverRequest(BaseModel): bolt11: str = Field(description="BOLT11 invoice string") +async def _request_mint_with_fallback( + amount_sats: int, +) -> tuple[str, str, str]: + """Primary first, fall back to other trusted mints on rate-limit/transport failure.""" + tried: list[str] = [] + candidates = [settings.primary_mint] + [ + m for m in settings.cashu_mints if m != settings.primary_mint + ] + for mint_url in candidates: + try: + wallet = await get_wallet(mint_url, "sat") + quote = await _mint_operation( + lambda: wallet.request_mint(amount_sats), + op_name="request_mint_invoice", + mint_url=mint_url, + ) + return quote.request, quote.quote, mint_url + except Exception as e: + tried.append(f"{mint_url}: {type(e).__name__}") + if not is_mint_connection_error(e) and not _is_mint_rate_limited(e): + raise + logger.warning( + "request_mint failed, trying fallback mint", + extra={ + "failed_mint": mint_url, + "error": str(e), + "tried": tried, + }, + ) + continue + raise MintConnectionError(f"All mints failed for request_mint: {tried}") + + async def generate_lightning_invoice( amount_sats: int, description: str -) -> tuple[str, str]: - wallet = await get_wallet(settings.primary_mint, "sat") - quote = await wallet.request_mint(amount_sats) - return quote.request, quote.quote +) -> tuple[str, str, str]: + bolt11, payment_hash, mint_url = await _request_mint_with_fallback(amount_sats) + return bolt11, payment_hash, mint_url def generate_invoice_id() -> str: @@ -99,7 +137,7 @@ async def create_invoice( try: description = f"Routstr {request.purpose} {request.amount_sats} sats" - bolt11, payment_hash = await generate_lightning_invoice( + bolt11, payment_hash, mint_url = await generate_lightning_invoice( request.amount_sats, description ) @@ -115,6 +153,7 @@ async def create_invoice( status="pending", api_key_hash=api_key_token[3:] if api_key_token else None, purpose=request.purpose, + mint_url=mint_url, balance_limit=request.balance_limit, balance_limit_reset=request.balance_limit_reset, validity_date=request.validity_date, @@ -223,9 +262,14 @@ async def check_invoice_payment( invoice: LightningInvoice, session: AsyncSession ) -> None: try: - wallet = await get_wallet(settings.primary_mint, "sat") + mint_url = invoice.mint_url or settings.primary_mint + wallet = await get_wallet(mint_url, "sat") - mint_status = await wallet.get_mint_quote(invoice.payment_hash) + mint_status = await _mint_operation( + lambda: wallet.get_mint_quote(invoice.payment_hash), + op_name="get_mint_quote", + mint_url=mint_url, + ) if mint_status.paid: invoice.status = "paid" @@ -258,8 +302,13 @@ async def check_invoice_payment( async def create_api_key_from_invoice( invoice: LightningInvoice, session: AsyncSession ) -> ApiKey: - wallet = await get_wallet(settings.primary_mint, "sat") - await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) + mint_url = invoice.mint_url or settings.primary_mint + wallet = await get_wallet(mint_url, "sat") + await _mint_operation( + lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), + op_name="invoice_mint_create", + mint_url=mint_url, + ) dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest() @@ -268,7 +317,7 @@ async def create_api_key_from_invoice( hashed_key=hashed_key, balance=invoice.amount_sats * 1000, # Convert to msats refund_currency="sat", - refund_mint_url=settings.primary_mint, + refund_mint_url=mint_url, balance_limit=invoice.balance_limit, balance_limit_reset=invoice.balance_limit_reset, validity_date=invoice.validity_date, @@ -283,8 +332,13 @@ async def create_api_key_from_invoice( async def topup_api_key_from_invoice( invoice: LightningInvoice, session: AsyncSession ) -> None: - wallet = await get_wallet(settings.primary_mint, "sat") - await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) + mint_url = invoice.mint_url or settings.primary_mint + wallet = await get_wallet(mint_url, "sat") + await _mint_operation( + lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), + op_name="invoice_mint_topup", + mint_url=mint_url, + ) if not invoice.api_key_hash: raise ValueError("No API key associated with topup invoice") @@ -297,7 +351,9 @@ async def topup_api_key_from_invoice( await session.flush() -INVOICE_WATCH_INTERVAL_SECONDS = 5 +# Nutshell mints throttle Lightning backend lookups to once per 10s per +# quote, so polling faster just burns the global request budget for nothing. +INVOICE_WATCH_INTERVAL_SECONDS = 10 INVOICE_WATCH_BATCH_LIMIT = 100 diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index 26cf580d..3e67412d 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -1,11 +1,15 @@ from __future__ import annotations +import asyncio import math from typing import TypedDict import httpx from cashu.wallet.wallet import Proof, Wallet +from ..core.settings import settings +from ..wallet import _mint_operation + try: from bech32 import bech32_decode, convertbits # type: ignore except ModuleNotFoundError: # pragma: no cover – allow runtime miss @@ -215,15 +219,23 @@ async def raw_send_to_lnurl( lnurl_data["callback_url"], final_amount ) - melt_quote_resp = await wallet.melt_quote(invoice=bolt11_invoice) + melt_quote_resp = await _mint_operation( + lambda: wallet.melt_quote(invoice=bolt11_invoice), + op_name="lnurl_melt_quote", + mint_url=str(wallet.url), + ) if amount: proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) - _ = await wallet.melt( - proofs=proofs, - invoice=bolt11_invoice, - fee_reserve_sat=melt_quote_resp.fee_reserve, - quote_id=melt_quote_resp.quote, + _ = await _mint_operation( + lambda: wallet.melt( + proofs=proofs, + invoice=bolt11_invoice, + fee_reserve_sat=melt_quote_resp.fee_reserve, + quote_id=melt_quote_resp.quote, + ), + op_name="lnurl_melt", + mint_url=str(wallet.url), ) return final_amount diff --git a/routstr/wallet.py b/routstr/wallet.py index ccf7877f..6de0a6c9 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -16,7 +16,6 @@ from sqlmodel import col, select, update from .core import db, get_logger from .core.db import store_cashu_transaction from .core.settings import settings -from .payment.lnurl import raw_send_to_lnurl # cashu still declares Optional[X] without explicit defaults on MintInfo. # Under pydantic v2 those are required, but real mints omit many of them. @@ -62,6 +61,163 @@ _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( ) +class _MintRateLimiter: + """Per-mint token-bucket rate limiter. + + Enforces a maximum number of mint API requests per minute per mint URL. + When the bucket is empty, callers block until a token is available. + This prevents the node runner from being rate-limited or blocked by + mints that enforce request quotas. + """ + + _limiters: dict[str, "_MintRateLimiter"] = {} + + @classmethod + def get(cls, mint_url: str) -> "_MintRateLimiter | None": + rpm = settings.mint_max_requests_per_minute + if rpm <= 0: + return None + if mint_url not in cls._limiters: + cls._limiters[mint_url] = cls(mint_url, rpm) + return cls._limiters[mint_url] + + def __init__(self, mint_url: str, max_per_minute: int): + self._mint_url = mint_url + self._max = max_per_minute + # Refill rate: tokens per second + self._refill_rate = max_per_minute / 60.0 + self._tokens: float = float(max_per_minute) + self._last_refill = time.monotonic() + self._lock = asyncio.Lock() + + async def acquire(self) -> None: + async with self._lock: + now = time.monotonic() + elapsed = now - self._last_refill + self._tokens = min(self._max, self._tokens + elapsed * self._refill_rate) + self._last_refill = now + if self._tokens < 1: + wait = (1 - self._tokens) / self._refill_rate + logger.debug( + "Mint rate limiter: throttling", + extra={ + "mint_url": self._mint_url, + "wait_seconds": round(wait, 2), + "tokens_available": round(self._tokens, 2), + }, + ) + await asyncio.sleep(wait) + self._tokens = 0 + else: + self._tokens -= 1 + + +def _is_mint_rate_limited(error: BaseException) -> bool: + """True if the mint returned a 429 or rate-limit indication.""" + current: BaseException | None = error + seen: set[int] = set() + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, httpx.HTTPStatusError): + if current.response.status_code == 429: + return True + lowered = str(current).lower() + if "rate limit" in lowered or "too many requests" in lowered: + return True + current = current.__cause__ or current.__context__ + return False + + +async def _mint_operation( + factory, *, op_name: str = "mint_operation", mint_url: str = "" +): + """Wrap a mint API callable with rate limiting, timeout, and retry. + + ``factory`` must be a zero-arg callable that returns a fresh coroutine + each call — a pre-created coroutine can only be awaited once, so on retry + the original would be dead. + """ + limiter = _MintRateLimiter.get(mint_url) if mint_url else None + timeout = settings.mint_operation_timeout_seconds + max_attempts = settings.mint_retry_max_attempts + 1 + + last_exc: Exception | None = None + for attempt in range(max_attempts): + if limiter is not None: + await limiter.acquire() + + try: + if timeout > 0: + return await asyncio.wait_for(factory(), timeout=timeout) + return await factory() + except asyncio.TimeoutError as exc: + last_exc = exc + if attempt < max_attempts - 1: + backoff = (2 ** attempt) + (time.monotonic() % 1.0) + logger.warning( + "Mint operation timed out, retrying", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "backoff_seconds": round(backoff, 2), + }, + ) + await asyncio.sleep(backoff) + continue + raise httpx.TimeoutException( + f"{op_name} timed out after {timeout}s (retried {attempt + 1}x)" + ) from exc + except httpx.HTTPStatusError as exc: + if _is_mint_rate_limited(exc) and attempt < max_attempts - 1: + backoff = (2 ** attempt) + (time.monotonic() % 1.0) + retry_after = _parse_retry_after(exc.response.headers) + if retry_after is not None: + backoff = min(retry_after, backoff * 2) + logger.warning( + "Mint returned 429, backing off", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "backoff_seconds": round(backoff, 2), + }, + ) + await asyncio.sleep(backoff) + continue + raise + except Exception as exc: + if _is_mint_rate_limited(exc) and attempt < max_attempts - 1: + backoff = (2 ** attempt) + (time.monotonic() % 1.0) + logger.warning( + "Mint rate-limited, backing off", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "backoff_seconds": round(backoff, 2), + }, + ) + await asyncio.sleep(backoff) + continue + raise + + if last_exc: + raise last_exc + raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") + + +def _parse_retry_after(headers) -> float | None: + """Parse a Retry-After header (delta-seconds form) into seconds.""" + raw = headers.get("retry-after") or headers.get("Retry-After") + if raw is None: + return None + try: + return float(str(raw).strip()) + except (TypeError, ValueError): + return None + + def is_mint_connection_error(error: BaseException) -> bool: """True if ``error`` (or anything in its cause/context chain) is a mint transport failure. Walks the chain because some sites re-raise transport @@ -192,10 +348,18 @@ async def _redeem_same_mint( that, not the face value, or routstr over-credits the user and its wallet drifts insolvent. """ - await wallet.load_mint(keyset_id=token_obj.keysets[0]) + await _mint_operation( + lambda: wallet.load_mint(keyset_id=token_obj.keysets[0]), + op_name="redeem_load_mint", + mint_url=token_obj.mint, + ) wallet.verify_proofs_dleq(token_obj.proofs) input_fees = wallet.get_fees_for_proofs(token_obj.proofs) - await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True) + await _mint_operation( + lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), + op_name="redeem_split", + mint_url=token_obj.mint, + ) return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint @@ -341,6 +505,44 @@ def _melt_insufficient_shortfall(error: Exception) -> int | None: return 1 +async def _request_mint_with_fallback( + amount: int, *, op_name: str, primary_wallet: Wallet | None = None +) -> tuple[Wallet, str, object]: + """Try request_mint on the primary mint, fall back to other trusted mints + on transport or rate-limit failure. Returns the wallet, mint_url, and quote.""" + candidates = [settings.primary_mint] + [ + m for m in settings.cashu_mints if m != settings.primary_mint + ] + tried: list[str] = [] + for mint_url in candidates: + try: + if mint_url == settings.primary_mint and primary_wallet is not None: + wallet = primary_wallet + else: + wallet = await get_wallet(mint_url, settings.primary_mint_unit) + quote = await _mint_operation( + lambda: wallet.request_mint(amount), + op_name=op_name, + mint_url=mint_url, + ) + return wallet, mint_url, quote + except Exception as e: + tried.append(f"{mint_url}: {type(e).__name__}") + if not is_mint_connection_error(e) and not _is_mint_rate_limited(e): + raise + logger.warning( + "request_mint failed, trying fallback mint", + extra={ + "failed_mint": mint_url, + "error": str(e), + "tried": tried, + "op_name": op_name, + }, + ) + continue + raise MintConnectionError(f"All mints failed for {op_name}: {tried}") + + async def _calculate_swap_amount( amount_msat: int, token_unit: str, @@ -374,8 +576,16 @@ async def _calculate_swap_amount( ) try: - dummy_mint_quote = await primary_wallet.request_mint(receive_amount) - dummy_melt_quote = await token_wallet.melt_quote(dummy_mint_quote.request) + _, _, dummy_mint_quote = await _request_mint_with_fallback( + receive_amount, + op_name="swap_fee_est_mint_quote", + primary_wallet=primary_wallet, + ) + dummy_melt_quote = await _mint_operation( + lambda: token_wallet.melt_quote(dummy_mint_quote.request), + op_name="swap_fee_est_melt_quote", + mint_url=token_mint_url, + ) fee_reserve = dummy_melt_quote.fee_reserve input_fees = token_wallet.get_fees_for_proofs(proofs) @@ -462,15 +672,23 @@ async def swap_to_primary_mint( # amount recomputed from the fees the mint actually demands. observed_extra_fee = 0 attempt = 0 + dest_wallet = primary_wallet + dest_mint_url = settings.primary_mint while True: attempt += 1 - mint_quote = await primary_wallet.request_mint(minted_amount) + dest_wallet, dest_mint_url, mint_quote = await _request_mint_with_fallback( + minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet + ) logger.info( "swap_to_primary_mint: mint quote received", - extra={"mint_quote_id": mint_quote.quote, "attempt": attempt}, + extra={"mint_quote_id": mint_quote.quote, "attempt": attempt, "dest_mint": dest_mint_url}, ) - melt_quote = await token_wallet.melt_quote(mint_quote.request) + melt_quote = await _mint_operation( + lambda: token_wallet.melt_quote(mint_quote.request), + op_name="swap_melt_quote", + mint_url=token_obj.mint, + ) input_fees = token_wallet.get_fees_for_proofs(token_obj.proofs) total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees logger.info( @@ -523,11 +741,15 @@ async def swap_to_primary_mint( continue try: - _ = await token_wallet.melt( - proofs=token_obj.proofs, - invoice=mint_quote.request, - fee_reserve_sat=melt_quote.fee_reserve, - quote_id=melt_quote.quote, + _ = await _mint_operation( + lambda: token_wallet.melt( + proofs=token_obj.proofs, + invoice=mint_quote.request, + fee_reserve_sat=melt_quote.fee_reserve, + quote_id=melt_quote.quote, + ), + op_name="swap_melt", + mint_url=token_obj.mint, ) except Exception as e: # A down mint won't fix itself by retrying with a smaller amount. @@ -576,14 +798,18 @@ async def swap_to_primary_mint( break logger.info( - "swap_to_primary_mint: melt succeeded, minting on primary", - extra={"minted_amount": minted_amount, "mint_quote_id": mint_quote.quote}, + "swap_to_primary_mint: melt succeeded, minting on destination", + extra={"minted_amount": minted_amount, "mint_quote_id": mint_quote.quote, "dest_mint": dest_mint_url}, ) - await primary_wallet.load_proofs(reload=True) - pre_mint_balance = primary_wallet.available_balance.amount + await dest_wallet.load_proofs(reload=True) + pre_mint_balance = dest_wallet.available_balance.amount try: - _ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote) + _ = await _mint_operation( + lambda: dest_wallet.mint(minted_amount, quote_id=mint_quote.quote), + op_name="swap_mint_on_primary", + mint_url=dest_mint_url, + ) except Exception as e: if "11003" in str(e) or "outputs already signed" in str(e).lower(): # Previous mint call signed outputs at the mint but failed before @@ -594,10 +820,10 @@ async def swap_to_primary_mint( extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, ) try: - for keyset_id in primary_wallet.keysets: - await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) - await primary_wallet.load_proofs(reload=True) - post_recovery_balance = primary_wallet.available_balance.amount + for keyset_id in dest_wallet.keysets: + await dest_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) + await dest_wallet.load_proofs(reload=True) + post_recovery_balance = dest_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance logger.info( "swap_to_primary_mint: recovery scan completed", @@ -648,14 +874,14 @@ async def swap_to_primary_mint( "swap_to_primary_mint: completed successfully", extra={ "foreign_mint": token_obj.mint, - "primary_mint": settings.primary_mint, + "dest_mint": dest_mint_url, "original_amount": token_amount, "minted_amount": minted_amount, "unit": settings.primary_mint_unit, }, ) - return int(minted_amount), settings.primary_mint_unit, settings.primary_mint + return int(minted_amount), settings.primary_mint_unit, dest_mint_url async def credit_balance( @@ -760,17 +986,35 @@ async def credit_balance( _wallets: dict[str, Wallet] = {} +_wallet_last_load: dict[str, float] = {} +# Minimum seconds between full mint info + proof reloads for the same +# wallet. Prevents redundant mint API calls when get_wallet(load=True) +# is called rapidly by multiple background tasks (balance fetch, payout, +# auto-topup all hitting get_wallet within the same cycle). +_WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30 async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet: - global _wallets + global _wallets, _wallet_last_load id = f"{mint_url}_{unit}" if id not in _wallets: _wallets[id] = await Wallet.with_db(mint_url, db=".wallet", unit=unit) if load: - await _wallets[id].load_mint() - await _wallets[id].load_proofs(reload=True) + now = time.monotonic() + last = _wallet_last_load.get(id, 0) + if now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: + await _mint_operation( + lambda: _wallets[id].load_mint(), + op_name="load_mint", + mint_url=mint_url, + ) + await _mint_operation( + lambda: _wallets[id].load_proofs(reload=True), + op_name="load_proofs", + mint_url=mint_url, + ) + _wallet_last_load[id] = now return _wallets[id] @@ -788,20 +1032,35 @@ def get_proofs_per_mint_and_unit( return proofs -async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[Proof]: +async def slow_filter_spend_proofs( + proofs: list[Proof], wallet: Wallet +) -> list[Proof]: if not proofs: return [] _proofs = [] _spent_proofs = [] - for i in range(0, len(proofs), 1000): - pb = proofs[i : i + 1000] - proof_states = await wallet.check_proof_state(pb) + # Smaller batch size to reduce per-request load on the mint. + # 1000 proofs per batch was too aggressive and triggered rate limits + # on mints with strict request quotas. + batch_size = 100 + for i in range(0, len(proofs), batch_size): + pb = proofs[i : i + batch_size] + proof_states = await _mint_operation( + lambda: wallet.check_proof_state(pb), + op_name="check_proof_state", + mint_url=str(wallet.url), + ) for proof, state in zip(pb, proof_states.states): if str(state.state) != "spent": _proofs.append(proof) else: _spent_proofs.append(proof) - await wallet.set_reserved_for_send(_spent_proofs, reserved=True) + if _spent_proofs: + await _mint_operation( + lambda: wallet.set_reserved_for_send(_spent_proofs, reserved=True), + op_name="set_reserved_spent_proofs", + mint_url=str(wallet.url), + ) return _proofs @@ -923,6 +1182,8 @@ async def periodic_payout() -> None: if not settings.receive_ln_address: continue try: + from .payment.lnurl import raw_send_to_lnurl + async with db.create_session() as session: for mint_url in settings.cashu_mints: for unit in ["sat", "msat"]: @@ -1037,6 +1298,8 @@ async def periodic_routstr_fee_payout() -> None: while True: await asyncio.sleep(ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS) try: + from .payment.lnurl import raw_send_to_lnurl + async with db.create_session() as session: fee = await db.get_routstr_fee(session) accumulated_sats = fee.accumulated_msats // 1000 @@ -1065,6 +1328,8 @@ async def periodic_routstr_fee_payout() -> None: async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: + from .payment.lnurl import raw_send_to_lnurl + wallet = await get_wallet(mint, unit) proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) diff --git a/tests/integration/test_lightning_invoice_rip08.py b/tests/integration/test_lightning_invoice_rip08.py index 29301a42..35f1f1e7 100644 --- a/tests/integration/test_lightning_invoice_rip08.py +++ b/tests/integration/test_lightning_invoice_rip08.py @@ -26,11 +26,12 @@ async def patch_invoice_generation() -> Any: """Stub out `generate_lightning_invoice` so no mint round-trip is needed.""" counter = {"n": 0} - async def fake_generate(amount_sats: int, description: str) -> tuple[str, str]: + async def fake_generate(amount_sats: int, description: str) -> tuple[str, str, str]: counter["n"] += 1 return ( f"lnbc{amount_sats}n1pfakeinvoice{counter['n']}", f"payment_hash_{counter['n']}", + "http://localhost:3338", ) with patch( diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 3bb36a28..cc35e63a 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1321,3 +1321,244 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: await swap_to_primary_mint(mock_token, mock_token_wallet) assert mock_token_wallet.melt.call_count == 1 + + +# --------------------------------------------------------------------------- +# _mint_operation factory + retry +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_mint_operation_factory_retry_succeeds() -> None: + """_mint_operation accepts a zero-arg factory, not a dead coroutine. + A factory that raises twice then succeeds must be retried and return.""" + from routstr.core.settings import settings + from routstr.wallet import _mint_operation + + calls = 0 + + async def factory(): + nonlocal calls + calls += 1 + if calls < 3: + raise TimeoutError("timeout") + return "ok" + + with patch.object(settings, "mint_retry_max_attempts", 3): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch("asyncio.sleep", AsyncMock()): + result = await _mint_operation( + factory, op_name="test_retry", mint_url="http://mint:3338" + ) + + assert calls == 3 + assert result == "ok" + + +@pytest.mark.asyncio +async def test_mint_operation_factory_retry_exhausted() -> None: + """When the factory always times out, _mint_operation raises + httpx.TimeoutException after mint_retry_max_attempts + 1 attempts.""" + from routstr.core.settings import settings + from routstr.wallet import _mint_operation + + calls = 0 + + async def factory(): + nonlocal calls + calls += 1 + raise TimeoutError("always timeout") + + with patch.object(settings, "mint_retry_max_attempts", 2): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch("asyncio.sleep", AsyncMock()): + with pytest.raises(httpx.TimeoutException): + await _mint_operation( + factory, op_name="test_exhaust", mint_url="http://mint:3338" + ) + + assert calls == 3 # max_attempts(2) + 1 + + +# --------------------------------------------------------------------------- +# Trusted-mint fallback +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_lightning_mint_fallback_for_topups() -> None: + """When the primary mint is unreachable, _request_mint_with_fallback + falls back to a secondary trusted mint.""" + from routstr.core.settings import settings + from routstr.lightning import _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock( + side_effect=httpx.ConnectError("primary down") + ) + + mock_quote = Mock() + mock_quote.request = "lnbc1secondary" + mock_quote.quote = "quote_secondary" + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("routstr.lightning.get_wallet", side_effect=mock_get): + bolt11, quote_id, mint_url = await _request_mint_with_fallback( + 1000 + ) + + assert mint_url == secondary + assert bolt11 == "lnbc1secondary" + assert quote_id == "quote_secondary" + mock_primary_wallet.request_mint.assert_called_once() + mock_secondary_wallet.request_mint.assert_called_once() + + +@pytest.mark.asyncio +async def test_swap_falls_back_to_secondary_mint() -> None: + """When the primary mint is unreachable, swap_to_primary_mint falls back + to a secondary trusted mint as the swap destination.""" + from routstr.core.settings import settings + from routstr.wallet import _wallet_last_load, _wallets, swap_to_primary_mint + + _wallets.clear() + _wallet_last_load.clear() + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + foreign = "http://foreign:3338" + + mock_token = Mock() + mock_token.mint = foreign + mock_token.unit = "sat" + mock_token.amount = 1000 + mock_token.keysets = ["keyset1"] + mock_token.proofs = [Mock(amount=1000)] + + mock_token_wallet = Mock() + mock_token_wallet.load_mint = AsyncMock() + mock_token_wallet.load_proofs = AsyncMock() + mock_token_wallet.get_fees_for_proofs = Mock(return_value=0) + mock_token_wallet.melt_quote = AsyncMock( + return_value=Mock(quote="melt_q", amount=990, fee_reserve=10) + ) + mock_token_wallet.melt = AsyncMock(return_value=Mock()) + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock( + side_effect=httpx.ConnectError("primary down") + ) + + mint_quote = Mock(quote="mint_q_secondary", request="lnbc1secondary") + mock_secondary_wallet = Mock() + mock_secondary_wallet.load_mint = AsyncMock() + mock_secondary_wallet.load_proofs = AsyncMock() + mock_secondary_wallet.available_balance = Mock(amount=0) + mock_secondary_wallet.keysets = ["ks_secondary"] + mock_secondary_wallet.restore_tokens_for_keyset = AsyncMock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mint_quote) + mock_secondary_wallet.mint = AsyncMock(return_value=Mock()) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "primary_mint_unit", "sat"): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("asyncio.sleep", AsyncMock()): + with patch( + "routstr.wallet.get_wallet", side_effect=mock_get + ): + amount, unit, mint_url = ( + await swap_to_primary_mint( + mock_token, mock_token_wallet + ) + ) + + assert mint_url == secondary + assert amount == 990 # 1000 - 10 fee_reserve + assert unit == "sat" + mock_secondary_wallet.mint.assert_called_once() + mock_primary_wallet.mint.assert_not_called() + + +@pytest.mark.asyncio +async def test_lightning_mint_fallback_on_429() -> None: + """A 429 from the primary mint should trigger fallback to a secondary, + not just transport errors.""" + from routstr.core.settings import settings + from routstr.lightning import _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + mock_resp = Mock(status_code=429, headers={}) + mock_resp.raise_for_status = Mock(side_effect=httpx.HTTPStatusError( + "rate limited", request=Mock(), response=mock_resp + )) + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock( + side_effect=httpx.HTTPStatusError("rate limited", request=Mock(), response=mock_resp) + ) + + mock_quote = Mock(request="lnbc1secondary", quote="quote_secondary") + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_retry_max_attempts", 0): + with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("routstr.lightning.get_wallet", side_effect=mock_get): + bolt11, quote_id, mint_url = await _request_mint_with_fallback(1000) + + assert mint_url == secondary + mock_secondary_wallet.request_mint.assert_called_once() + + +@pytest.mark.asyncio +async def test_lightning_mint_fallback_all_fail() -> None: + """When every trusted mint fails, _request_mint_with_fallback raises + MintConnectionError instead of trying indefinitely.""" + from routstr.core.settings import settings + from routstr.lightning import _request_mint_with_fallback + from routstr.wallet import MintConnectionError + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock(side_effect=httpx.ConnectError("down")) + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(side_effect=httpx.ConnectError("down")) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_retry_max_attempts", 0): + with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("routstr.lightning.get_wallet", side_effect=mock_get): + with pytest.raises(MintConnectionError): + await _request_mint_with_fallback(1000) From d8db2a3051ae2323c86e4f2729a9edaf335de8b3 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 10 Jul 2026 21:46:56 +0200 Subject: [PATCH 04/46] fix: harden mint rate limiting and fallback --- routstr/lightning.py | 34 ++- routstr/payment/lnurl.py | 1 + routstr/wallet.py | 131 +++++++----- .../test_lightning_invoice_rip08.py | 10 +- tests/unit/test_wallet.py | 200 +++++++++++++++--- 5 files changed, 286 insertions(+), 90 deletions(-) diff --git a/routstr/lightning.py b/routstr/lightning.py index 24c21dde..c7f0919b 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -72,12 +72,13 @@ class InvoiceRecoverRequest(BaseModel): async def _request_mint_with_fallback( amount_sats: int, + *, + allowed_mints: list[str] | None = None, ) -> tuple[str, str, str]: - """Primary first, fall back to other trusted mints on rate-limit/transport failure.""" + """Request a quote, falling back only among the allowed trusted mints.""" tried: list[str] = [] - candidates = [settings.primary_mint] + [ - m for m in settings.cashu_mints if m != settings.primary_mint - ] + configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints] + candidates = list(dict.fromkeys(configured)) for mint_url in candidates: try: wallet = await get_wallet(mint_url, "sat") @@ -104,9 +105,14 @@ async def _request_mint_with_fallback( async def generate_lightning_invoice( - amount_sats: int, description: str + amount_sats: int, + description: str, + *, + allowed_mints: list[str] | None = None, ) -> tuple[str, str, str]: - bolt11, payment_hash, mint_url = await _request_mint_with_fallback(amount_sats) + bolt11, payment_hash, mint_url = await _request_mint_with_fallback( + amount_sats, allowed_mints=allowed_mints + ) return bolt11, payment_hash, mint_url @@ -121,6 +127,7 @@ async def create_invoice( session: AsyncSession = Depends(get_session), ) -> InvoiceCreateResponse: api_key_token = _extract_bearer_api_key(authorization) or request.api_key + topup_api_key: ApiKey | None = None if request.purpose == "topup": if not api_key_token: @@ -131,14 +138,21 @@ async def create_invoice( if not api_key_token.startswith("sk-"): raise HTTPException(status_code=400, detail="Invalid API key format") - api_key = await session.get(ApiKey, api_key_token[3:]) - if not api_key: + topup_api_key = await session.get(ApiKey, api_key_token[3:]) + if not topup_api_key: raise HTTPException(status_code=404, detail="API key not found") try: description = f"Routstr {request.purpose} {request.amount_sats} sats" + # An API key is backed by one mint. A top-up must use that same mint; + # falling back to another would create mixed-mint collateral that the + # current single refund_mint_url field cannot account for or refund. + allowed_mints = None + if request.purpose == "topup": + assert topup_api_key is not None + allowed_mints = [topup_api_key.refund_mint_url or settings.primary_mint] bolt11, payment_hash, mint_url = await generate_lightning_invoice( - request.amount_sats, description + request.amount_sats, description, allowed_mints=allowed_mints ) invoice_id = generate_invoice_id() @@ -308,6 +322,7 @@ async def create_api_key_from_invoice( lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), op_name="invoice_mint_create", mint_url=mint_url, + retry_timeouts=False, ) dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" @@ -338,6 +353,7 @@ async def topup_api_key_from_invoice( lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), op_name="invoice_mint_topup", mint_url=mint_url, + retry_timeouts=False, ) if not invoice.api_key_hash: diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index 1ad5e06a..fbd28586 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -244,6 +244,7 @@ async def raw_send_to_lnurl( ), op_name="lnurl_melt", mint_url=str(wallet.url), + retry_timeouts=False, ), timeout=MELT_TIMEOUT_SECONDS, ) diff --git a/routstr/wallet.py b/routstr/wallet.py index d8f5241f..a5a4e8cd 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -78,9 +78,11 @@ class _MintRateLimiter: rpm = settings.mint_max_requests_per_minute if rpm <= 0: return None - if mint_url not in cls._limiters: - cls._limiters[mint_url] = cls(mint_url, rpm) - return cls._limiters[mint_url] + limiter = cls._limiters.get(mint_url) + if limiter is None or limiter._max != rpm: + limiter = cls(mint_url, rpm) + cls._limiters[mint_url] = limiter + return limiter def __init__(self, mint_url: str, max_per_minute: int): self._mint_url = mint_url @@ -93,11 +95,17 @@ class _MintRateLimiter: async def acquire(self) -> None: async with self._lock: - now = time.monotonic() - elapsed = now - self._last_refill - self._tokens = min(self._max, self._tokens + elapsed * self._refill_rate) - self._last_refill = now - if self._tokens < 1: + while True: + now = time.monotonic() + elapsed = now - self._last_refill + self._tokens = min( + self._max, self._tokens + elapsed * self._refill_rate + ) + self._last_refill = now + if self._tokens >= 1: + self._tokens -= 1 + return + wait = (1 - self._tokens) / self._refill_rate logger.debug( "Mint rate limiter: throttling", @@ -108,9 +116,6 @@ class _MintRateLimiter: }, ) await asyncio.sleep(wait) - self._tokens = 0 - else: - self._tokens -= 1 def _is_mint_rate_limited(error: BaseException) -> bool: @@ -130,7 +135,11 @@ def _is_mint_rate_limited(error: BaseException) -> bool: async def _mint_operation( - factory: Callable[[], Awaitable[Any]], *, op_name: str = "mint_operation", mint_url: str = "" + factory: Callable[[], Awaitable[Any]], + *, + op_name: str = "mint_operation", + mint_url: str = "", + retry_timeouts: bool = True, ) -> Any: """Wrap a mint API callable with rate limiting, timeout, and retry. @@ -151,10 +160,10 @@ async def _mint_operation( if timeout > 0: return await asyncio.wait_for(factory(), timeout=timeout) return await factory() - except asyncio.TimeoutError as exc: + except (asyncio.TimeoutError, httpx.TimeoutException) as exc: last_exc = exc - if attempt < max_attempts - 1: - backoff = (2 ** attempt) + (time.monotonic() % 1.0) + if retry_timeouts and attempt < max_attempts - 1: + backoff = (2**attempt) + (time.monotonic() % 1.0) logger.warning( "Mint operation timed out, retrying", extra={ @@ -171,10 +180,10 @@ async def _mint_operation( ) from exc except httpx.HTTPStatusError as exc: if _is_mint_rate_limited(exc) and attempt < max_attempts - 1: - backoff = (2 ** attempt) + (time.monotonic() % 1.0) + backoff = (2**attempt) + (time.monotonic() % 1.0) retry_after = _parse_retry_after(exc.response.headers) if retry_after is not None: - backoff = min(retry_after, backoff * 2) + backoff = max(retry_after, backoff) logger.warning( "Mint returned 429, backing off", extra={ @@ -189,7 +198,7 @@ async def _mint_operation( raise except Exception as exc: if _is_mint_rate_limited(exc) and attempt < max_attempts - 1: - backoff = (2 ** attempt) + (time.monotonic() % 1.0) + backoff = (2**attempt) + (time.monotonic() % 1.0) logger.warning( "Mint rate-limited, backing off", extra={ @@ -360,6 +369,7 @@ async def _redeem_same_mint( lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), op_name="redeem_split", mint_url=token_obj.mint, + retry_timeouts=False, ) return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint @@ -402,7 +412,9 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { - f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id) + f"{k.mint_url}/{k.unit.name}": sum( + p.amount for p in wallet.proofs if p.id == k.id + ) for k in wallet.keysets.values() } # Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet @@ -682,7 +694,11 @@ async def swap_to_primary_mint( ) logger.info( "swap_to_primary_mint: mint quote received", - extra={"mint_quote_id": mint_quote.quote, "attempt": attempt, "dest_mint": dest_mint_url}, + extra={ + "mint_quote_id": mint_quote.quote, + "attempt": attempt, + "dest_mint": dest_mint_url, + }, ) melt_quote = await _mint_operation( @@ -751,6 +767,7 @@ async def swap_to_primary_mint( ), op_name="swap_melt", mint_url=token_obj.mint, + retry_timeouts=False, ) except Exception as e: # A down mint won't fix itself by retrying with a smaller amount. @@ -800,7 +817,11 @@ async def swap_to_primary_mint( logger.info( "swap_to_primary_mint: melt succeeded, minting on destination", - extra={"minted_amount": minted_amount, "mint_quote_id": mint_quote.quote, "dest_mint": dest_mint_url}, + extra={ + "minted_amount": minted_amount, + "mint_quote_id": mint_quote.quote, + "dest_mint": dest_mint_url, + }, ) await dest_wallet.load_proofs(reload=True) @@ -810,6 +831,7 @@ async def swap_to_primary_mint( lambda: dest_wallet.mint(minted_amount, quote_id=mint_quote.quote), op_name="swap_mint_on_primary", mint_url=dest_mint_url, + retry_timeouts=False, ) except Exception as e: if "11003" in str(e) or "outputs already signed" in str(e).lower(): @@ -818,11 +840,16 @@ async def swap_to_primary_mint( # advance the counter so the next request derives fresh secrets. logger.warning( "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", - extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, + extra={ + "mint_quote_id": mint_quote.quote, + "minted_amount": minted_amount, + }, ) try: for keyset_id in dest_wallet.keysets: - await dest_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) + await dest_wallet.restore_tokens_for_keyset( + keyset_id, to=1, batch=25 + ) await dest_wallet.load_proofs(reload=True) post_recovery_balance = dest_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance @@ -988,6 +1015,7 @@ async def credit_balance( _wallets: dict[str, Wallet] = {} _wallet_last_load: dict[str, float] = {} +_wallet_load_locks: dict[str, asyncio.Lock] = {} # Minimum seconds between full mint info + proof reloads for the same # wallet. Prevents redundant mint API calls when get_wallet(load=True) # is called rapidly by multiple background tasks (balance fetch, payout, @@ -996,27 +1024,29 @@ _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30 async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet: - global _wallets, _wallet_last_load + global _wallets, _wallet_last_load, _wallet_load_locks id = f"{mint_url}_{unit}" - if id not in _wallets: - _wallets[id] = await Wallet.with_db(mint_url, db=".wallet", unit=unit) + lock = _wallet_load_locks.setdefault(id, asyncio.Lock()) + async with lock: + if id not in _wallets: + _wallets[id] = await Wallet.with_db(mint_url, db=".wallet", unit=unit) - if load: - now = time.monotonic() - last = _wallet_last_load.get(id, 0) - if now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: - await _mint_operation( - lambda: _wallets[id].load_mint(), - op_name="load_mint", - mint_url=mint_url, - ) - await _mint_operation( - lambda: _wallets[id].load_proofs(reload=True), - op_name="load_proofs", - mint_url=mint_url, - ) - _wallet_last_load[id] = now - return _wallets[id] + if load: + now = time.monotonic() + last = _wallet_last_load.get(id, 0) + if now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: + await _mint_operation( + lambda: _wallets[id].load_mint(), + op_name="load_mint", + mint_url=mint_url, + ) + await _mint_operation( + lambda: _wallets[id].load_proofs(reload=True), + op_name="load_proofs", + mint_url=mint_url, + ) + _wallet_last_load[id] = time.monotonic() + return _wallets[id] def get_proofs_per_mint_and_unit( @@ -1033,9 +1063,7 @@ def get_proofs_per_mint_and_unit( return proofs -async def slow_filter_spend_proofs( - proofs: list[Proof], wallet: Wallet -) -> list[Proof]: +async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[Proof]: if not proofs: return [] _proofs = [] @@ -1061,6 +1089,7 @@ async def slow_filter_spend_proofs( lambda: wallet.set_reserved_for_send(_spent_proofs, reserved=True), op_name="set_reserved_spent_proofs", mint_url=str(wallet.url), + retry_timeouts=False, ) return _proofs @@ -1108,7 +1137,9 @@ async def fetch_all_balances( "unit": unit, "wallet_balance": proofs_balance, "user_balance": user_balance, - "owner_balance": proofs_balance - user_balance if proofs_balance != 0 else 0, + "owner_balance": proofs_balance - user_balance + if proofs_balance != 0 + else 0, } return result except Exception as e: @@ -1316,7 +1347,6 @@ async def periodic_routstr_fee_payout() -> None: while True: await asyncio.sleep(ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS) try: - async with db.create_session() as session: fee = await db.get_routstr_fee(session) accumulated_sats = fee.accumulated_msats // 1000 @@ -1326,7 +1356,11 @@ async def periodic_routstr_fee_payout() -> None: wallet, settings.primary_mint, "sat", not_reserved=True ) amount_received = await raw_send_to_lnurl( - wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats + wallet, + proofs, + ROUTSTR_LN_ADDRESS, + "sat", + amount=accumulated_sats, ) paid_msats = accumulated_sats * 1000 await db.reset_routstr_fee(session, paid_msats) @@ -1345,7 +1379,6 @@ async def periodic_routstr_fee_payout() -> None: async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: - wallet = await get_wallet(mint, unit) proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) diff --git a/tests/integration/test_lightning_invoice_rip08.py b/tests/integration/test_lightning_invoice_rip08.py index 35f1f1e7..766d0176 100644 --- a/tests/integration/test_lightning_invoice_rip08.py +++ b/tests/integration/test_lightning_invoice_rip08.py @@ -26,7 +26,12 @@ async def patch_invoice_generation() -> Any: """Stub out `generate_lightning_invoice` so no mint round-trip is needed.""" counter = {"n": 0} - async def fake_generate(amount_sats: int, description: str) -> tuple[str, str, str]: + async def fake_generate( + amount_sats: int, + description: str, + *, + allowed_mints: list[str] | None = None, + ) -> tuple[str, str, str]: counter["n"] += 1 return ( f"lnbc{amount_sats}n1pfakeinvoice{counter['n']}", @@ -96,6 +101,9 @@ async def test_topup_with_authorization_header( body = resp.json() assert body["amount_sats"] == 500 assert body["bolt11"].startswith("lnbc") + assert patch_invoice_generation.call_args.kwargs["allowed_mints"] == [ + "http://localhost:3338" + ] @pytest.mark.integration diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 233fbd1f..d949050a 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1,3 +1,4 @@ +import asyncio import base64 import json import socket @@ -19,6 +20,26 @@ from routstr.wallet import ( ) +@pytest.fixture(autouse=True) +def isolate_wallet_runtime_state(): + """Keep production limiter/wallet caches from leaking across unit tests.""" + from routstr import wallet as wallet_module + from routstr.core.settings import settings + + original_rpm = settings.mint_max_requests_per_minute + settings.mint_max_requests_per_minute = 0 + wallet_module._MintRateLimiter._limiters.clear() + wallet_module._wallets.clear() + wallet_module._wallet_last_load.clear() + wallet_module._wallet_load_locks.clear() + yield + settings.mint_max_requests_per_minute = original_rpm + wallet_module._MintRateLimiter._limiters.clear() + wallet_module._wallets.clear() + wallet_module._wallet_last_load.clear() + wallet_module._wallet_load_locks.clear() + + @pytest.mark.asyncio async def test_get_balance() -> None: mock_wallet = Mock() @@ -728,9 +749,7 @@ async def test_calculate_swap_amount_same_mint_short_circuit() -> None: quotes are requested.""" from routstr.wallet import _calculate_swap_amount - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 1000, fee_reserves=[] - ) + _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(1000, fee_reserves=[]) from routstr.core.settings import settings @@ -755,9 +774,7 @@ async def test_calculate_swap_amount_msat_primary_unit() -> None: """With an msat primary mint the dummy quote and result stay in msats.""" from routstr.wallet import _calculate_swap_amount - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 179, fee_reserves=[2] - ) + _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(179, fee_reserves=[2]) from routstr.core.settings import settings @@ -805,12 +822,8 @@ async def test_calculate_swap_amount_wraps_estimation_failure() -> None: """Estimation infrastructure failures surface as a single clear ValueError.""" from routstr.wallet import _calculate_swap_amount - _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( - 179, fee_reserves=[] - ) - mock_primary_wallet.request_mint = AsyncMock( - side_effect=Exception("mint offline") - ) + _, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(179, fee_reserves=[]) + mock_primary_wallet.request_mint = AsyncMock(side_effect=Exception("mint offline")) from routstr.core.settings import settings @@ -1158,7 +1171,9 @@ def test_is_mint_connection_error_detects_transport_failures( ValueError("Invalid Cashu token"), # Mint answered with an error status — reachable, so NOT a connection error. httpx.HTTPStatusError( - "500", request=httpx.Request("POST", "http://m"), response=httpx.Response(500) + "500", + request=httpx.Request("POST", "http://m"), + response=httpx.Response(500), ), RuntimeError("some internal fault"), ], @@ -1213,7 +1228,9 @@ def test_classify_zero_value(error: ValueError) -> None: def test_classify_generic_valueerror_is_not_zero_value() -> None: """A generic wallet ValueError still falls to the generic bucket — the zero-value match must not over-trigger.""" - classified = classify_redemption_error(ValueError("some unexpected wallet condition")) + classified = classify_redemption_error( + ValueError("some unexpected wallet condition") + ) assert classified is not None type_, status, _msg, code = classified assert (type_, status, code) == ( @@ -1276,7 +1293,9 @@ async def test_credit_balance_db_transport_error_is_token_consumed() -> None: @pytest.mark.asyncio -async def test_swap_fee_estimation_transport_error_raises_mint_connection_error() -> None: +async def test_swap_fee_estimation_transport_error_raises_mint_connection_error() -> ( + None +): """A transport failure while estimating fees is surfaced as MintConnectionError (→ 503), not a generic fee ValueError (→ 422).""" from routstr.wallet import swap_to_primary_mint @@ -1308,9 +1327,7 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( 1000, fee_reserves=[10, 10] ) - mock_token_wallet.melt = AsyncMock( - side_effect=httpx.ConnectTimeout("timed out") - ) + mock_token_wallet.melt = AsyncMock(side_effect=httpx.ConnectTimeout("timed out")) from routstr.core.settings import settings @@ -1324,10 +1341,119 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: # --------------------------------------------------------------------------- -# _mint_operation factory + retry +# Per-mint limiter + _mint_operation factory/retry # --------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_mint_rate_limiter_serializes_waiters_after_refill() -> None: + from routstr.wallet import _MintRateLimiter + + limiter = _MintRateLimiter("http://mint:3338", 60) + limiter._tokens = 0 + limiter._last_refill = 0 + clock = {"now": 0.0} + sleeps: list[float] = [] + real_sleep = asyncio.sleep + + async def fake_sleep(delay: float) -> None: + sleeps.append(delay) + clock["now"] += delay + await real_sleep(0) + + with patch("routstr.wallet.time.monotonic", side_effect=lambda: clock["now"]): + with patch("routstr.wallet.asyncio.sleep", side_effect=fake_sleep): + await asyncio.gather(limiter.acquire(), limiter.acquire()) + + assert sleeps == pytest.approx([1.0, 1.0]) + assert clock["now"] == pytest.approx(2.0) + + +def test_mint_rate_limiter_rebuilds_when_setting_changes() -> None: + from routstr.core.settings import settings + from routstr.wallet import _MintRateLimiter + + with patch.object(settings, "mint_max_requests_per_minute", 20): + first = _MintRateLimiter.get("http://mint:3338") + with patch.object(settings, "mint_max_requests_per_minute", 10): + second = _MintRateLimiter.get("http://mint:3338") + + assert first is not None + assert second is not None + assert first is not second + assert second._max == 10 + + +@pytest.mark.asyncio +async def test_mint_operation_honors_retry_after_as_minimum() -> None: + from routstr.core.settings import settings + from routstr.wallet import _mint_operation + + request = httpx.Request("POST", "http://mint:3338/v1/mint/quote/bolt11") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + calls = 0 + + async def factory() -> str: + nonlocal calls + calls += 1 + if calls == 1: + raise httpx.HTTPStatusError( + "rate limited", request=request, response=response + ) + return "ok" + + sleep = AsyncMock() + with patch.object(settings, "mint_retry_max_attempts", 1): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("routstr.wallet.time.monotonic", return_value=0.1): + with patch("routstr.wallet.asyncio.sleep", sleep): + result = await _mint_operation(factory, mint_url="http://mint:3338") + + assert result == "ok" + sleep.assert_awaited_once_with(60.0) + + +@pytest.mark.asyncio +async def test_mint_operation_retries_httpx_timeout_only_when_safe() -> None: + from routstr.core.settings import settings + from routstr.wallet import _mint_operation + + retrying = AsyncMock(side_effect=[httpx.ReadTimeout("slow"), "ok"]) + non_retrying = AsyncMock(side_effect=httpx.ReadTimeout("ambiguous")) + + with patch.object(settings, "mint_retry_max_attempts", 2): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("routstr.wallet.asyncio.sleep", AsyncMock()): + assert await _mint_operation(retrying) == "ok" + with pytest.raises(httpx.TimeoutException): + await _mint_operation(non_retrying, retry_timeouts=False) + + assert retrying.await_count == 2 + assert non_retrying.await_count == 1 + + +@pytest.mark.asyncio +async def test_get_wallet_initializes_and_loads_once_concurrently() -> None: + from routstr.wallet import get_wallet + + mock_wallet = Mock() + mock_wallet.load_mint = AsyncMock() + mock_wallet.load_proofs = AsyncMock() + + with patch( + "routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet) + ) as create: + with patch("routstr.wallet.time.monotonic", return_value=100.0): + first, second = await asyncio.gather( + get_wallet("http://mint:3338"), get_wallet("http://mint:3338") + ) + + assert first is second is mock_wallet + create.assert_awaited_once() + mock_wallet.load_mint.assert_awaited_once() + mock_wallet.load_proofs.assert_awaited_once_with(reload=True) + + @pytest.mark.asyncio async def test_mint_operation_factory_retry_succeeds() -> None: """_mint_operation accepts a zero-arg factory, not a dead coroutine. @@ -1484,10 +1610,8 @@ async def test_swap_falls_back_to_secondary_mint() -> None: with patch( "routstr.wallet.get_wallet", side_effect=mock_get ): - amount, unit, mint_url = ( - await swap_to_primary_mint( - mock_token, mock_token_wallet - ) + amount, unit, mint_url = await swap_to_primary_mint( + mock_token, mock_token_wallet ) assert mint_url == secondary @@ -1508,12 +1632,16 @@ async def test_lightning_mint_fallback_on_429() -> None: secondary = "http://secondary:3338" mock_resp = Mock(status_code=429, headers={}) - mock_resp.raise_for_status = Mock(side_effect=httpx.HTTPStatusError( - "rate limited", request=Mock(), response=mock_resp - )) + mock_resp.raise_for_status = Mock( + side_effect=httpx.HTTPStatusError( + "rate limited", request=Mock(), response=mock_resp + ) + ) mock_primary_wallet = Mock() mock_primary_wallet.request_mint = AsyncMock( - side_effect=httpx.HTTPStatusError("rate limited", request=Mock(), response=mock_resp) + side_effect=httpx.HTTPStatusError( + "rate limited", request=Mock(), response=mock_resp + ) ) mock_quote = Mock(request="lnbc1secondary", quote="quote_secondary") @@ -1528,8 +1656,14 @@ async def test_lightning_mint_fallback_on_429() -> None: with patch.object(settings, "mint_retry_max_attempts", 0): with patch.object(settings, "mint_max_requests_per_minute", 0): with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch("routstr.lightning.get_wallet", side_effect=mock_get): - bolt11, quote_id, mint_url = await _request_mint_with_fallback(1000) + with patch( + "routstr.lightning.get_wallet", side_effect=mock_get + ): + ( + bolt11, + quote_id, + mint_url, + ) = await _request_mint_with_fallback(1000) assert mint_url == secondary mock_secondary_wallet.request_mint.assert_called_once() @@ -1549,7 +1683,9 @@ async def test_lightning_mint_fallback_all_fail() -> None: mock_primary_wallet = Mock() mock_primary_wallet.request_mint = AsyncMock(side_effect=httpx.ConnectError("down")) mock_secondary_wallet = Mock() - mock_secondary_wallet.request_mint = AsyncMock(side_effect=httpx.ConnectError("down")) + mock_secondary_wallet.request_mint = AsyncMock( + side_effect=httpx.ConnectError("down") + ) wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) @@ -1559,6 +1695,8 @@ async def test_lightning_mint_fallback_all_fail() -> None: with patch.object(settings, "mint_retry_max_attempts", 0): with patch.object(settings, "mint_max_requests_per_minute", 0): with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch("routstr.lightning.get_wallet", side_effect=mock_get): + with patch( + "routstr.lightning.get_wallet", side_effect=mock_get + ): with pytest.raises(MintConnectionError): await _request_mint_with_fallback(1000) From d23c90b939548eeb7541d49ee5f5d40384c00bb5 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 10 Jul 2026 21:50:50 +0200 Subject: [PATCH 05/46] fix: type wallet test fixture --- tests/unit/test_wallet.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index d949050a..24381f9a 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -2,6 +2,7 @@ import asyncio import base64 import json import socket +from collections.abc import Generator from unittest.mock import AsyncMock, Mock, patch import httpx @@ -21,7 +22,7 @@ from routstr.wallet import ( @pytest.fixture(autouse=True) -def isolate_wallet_runtime_state(): +def isolate_wallet_runtime_state() -> Generator[None, None, None]: """Keep production limiter/wallet caches from leaking across unit tests.""" from routstr import wallet as wallet_module from routstr.core.settings import settings From 1230d528de774bd39288c4ac2104fa31aebc4601 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 10 Jul 2026 23:46:09 +0200 Subject: [PATCH 06/46] fix: avoid rate limiting balance proof checks --- routstr/wallet.py | 7 +++---- tests/unit/test_wallet.py | 19 +++++++++++++++++++ 2 files changed, 22 insertions(+), 4 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index a5a4e8cd..972acabd 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -1068,10 +1068,9 @@ async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[ return [] _proofs = [] _spent_proofs = [] - # Smaller batch size to reduce per-request load on the mint. - # 1000 proofs per batch was too aggressive and triggered rate limits - # on mints with strict request quotas. - batch_size = 100 + # Keep proof-state checks in large batches. Mint quotas count HTTP requests, + # so smaller batches make balance reads slower and more likely to hit 429s. + batch_size = 1000 for i in range(0, len(proofs), batch_size): pb = proofs[i : i + batch_size] proof_states = await _mint_operation( diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 24381f9a..264bf26b 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1346,6 +1346,25 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: # --------------------------------------------------------------------------- +@pytest.mark.asyncio +async def test_balance_proof_check_uses_large_batches_to_avoid_rate_limit() -> None: + """Balance reads must not turn a few hundred proofs into many mint requests.""" + from routstr.wallet import slow_filter_spend_proofs + + proofs = [Mock() for _ in range(250)] + states = [Mock(state="UNSPENT") for _ in proofs] + wallet = Mock() + wallet.url = "http://mint:3338" + wallet.check_proof_state = AsyncMock(return_value=Mock(states=states)) + wallet.set_reserved_for_send = AsyncMock() + + result = await slow_filter_spend_proofs(proofs, wallet) + + assert result == proofs + wallet.check_proof_state.assert_awaited_once_with(proofs) + wallet.set_reserved_for_send.assert_not_awaited() + + @pytest.mark.asyncio async def test_mint_rate_limiter_serializes_waiters_after_refill() -> None: from routstr.wallet import _MintRateLimiter From acb630f6cf3229c41a61d82b3f2da10984924c6d Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 10 Jul 2026 23:54:05 +0200 Subject: [PATCH 07/46] refactor: adapt mint throttling to 429 responses --- routstr/core/settings.py | 30 +++--- routstr/wallet.py | 195 +++++++++++++++++--------------------- tests/unit/test_wallet.py | 110 +++++++++++++-------- 3 files changed, 175 insertions(+), 160 deletions(-) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 88797ff2..cee8f919 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -55,17 +55,12 @@ class Settings(BaseSettings): mint_operation_timeout_seconds: int = Field( default=30, gt=0, env="MINT_OPERATION_TIMEOUT_SECONDS" ) - # Maximum mint API requests per minute, per mint URL. Nutshell mints - # (e.g. Minibits) enforce 20/min/IP on transaction endpoints (mint, melt, - # swap, quotes) and 60/min/IP globally. 20 stays under the transaction - # bucket since most calls here are transaction ops. 0 = unlimited. - mint_max_requests_per_minute: int = Field( - default=20, ge=0, env="MINT_MAX_REQUESTS_PER_MINUTE" - ) + # Maximum concurrent API operations per mint. Actual mint quotas vary by + # endpoint, so 429 responses drive adaptive cooldown instead of fixed RPM + # pacing. 0 = unlimited concurrency. + mint_max_concurrency: int = Field(default=4, ge=0, env="MINT_MAX_CONCURRENCY") # Max retries when a mint returns 429 or times out (exponential backoff). - mint_retry_max_attempts: int = Field( - default=3, ge=0, env="MINT_RETRY_MAX_ATTEMPTS" - ) + mint_retry_max_attempts: int = Field(default=3, ge=0, env="MINT_RETRY_MAX_ATTEMPTS") # Pricing # Default behavior: derive pricing from MODELS @@ -114,7 +109,9 @@ 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") 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") @@ -133,9 +130,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.""" @@ -298,7 +294,11 @@ class SettingsService: valid_fields = set(env_resolved.dict().keys()) merged_dict: dict[str, Any] = dict(env_resolved.dict()) merged_dict.update( - {k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields} + { + k: v + for k, v in db_json.items() + if v not in (None, "", [], {}) and k in valid_fields + } ) merged_dict = Settings(**merged_dict).dict() diff --git a/routstr/wallet.py b/routstr/wallet.py index 972acabd..273d1aaf 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -62,60 +62,46 @@ _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( ) -class _MintRateLimiter: - """Per-mint token-bucket rate limiter. +class _MintRateGuard: + """Bound concurrency and adapt to actual per-mint 429 responses.""" - Enforces a maximum number of mint API requests per minute per mint URL. - When the bucket is empty, callers block until a token is available. - This prevents the node runner from being rate-limited or blocked by - mints that enforce request quotas. - """ - - _limiters: dict[str, "_MintRateLimiter"] = {} + _guards: dict[str, "_MintRateGuard"] = {} @classmethod - def get(cls, mint_url: str) -> "_MintRateLimiter | None": - rpm = settings.mint_max_requests_per_minute - if rpm <= 0: + def get(cls, mint_url: str) -> "_MintRateGuard | None": + concurrency = settings.mint_max_concurrency + if concurrency <= 0: return None - limiter = cls._limiters.get(mint_url) - if limiter is None or limiter._max != rpm: - limiter = cls(mint_url, rpm) - cls._limiters[mint_url] = limiter - return limiter + guard = cls._guards.get(mint_url) + if guard is None or guard._max_concurrency != concurrency: + guard = cls(mint_url, concurrency) + cls._guards[mint_url] = guard + return guard - def __init__(self, mint_url: str, max_per_minute: int): + def __init__(self, mint_url: str, max_concurrency: int): self._mint_url = mint_url - self._max = max_per_minute - # Refill rate: tokens per second - self._refill_rate = max_per_minute / 60.0 - self._tokens: float = float(max_per_minute) - self._last_refill = time.monotonic() - self._lock = asyncio.Lock() + self._max_concurrency = max_concurrency + self._semaphore = asyncio.Semaphore(max_concurrency) + self._cooldown_until = 0.0 - async def acquire(self) -> None: - async with self._lock: - while True: - now = time.monotonic() - elapsed = now - self._last_refill - self._tokens = min( - self._max, self._tokens + elapsed * self._refill_rate - ) - self._last_refill = now - if self._tokens >= 1: - self._tokens -= 1 - return + def apply_cooldown(self, delay: float) -> None: + self._cooldown_until = max( + self._cooldown_until, time.monotonic() + max(0.0, delay) + ) - wait = (1 - self._tokens) / self._refill_rate + async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any: + async with self._semaphore: + wait = self._cooldown_until - time.monotonic() + if wait > 0: logger.debug( - "Mint rate limiter: throttling", + "Mint rate guard: cooling down", extra={ "mint_url": self._mint_url, "wait_seconds": round(wait, 2), - "tokens_available": round(self._tokens, 2), }, ) await asyncio.sleep(wait) + return await factory() def _is_mint_rate_limited(error: BaseException) -> bool: @@ -141,80 +127,77 @@ async def _mint_operation( mint_url: str = "", retry_timeouts: bool = True, ) -> Any: - """Wrap a mint API callable with rate limiting, timeout, and retry. + """Run a mint operation with bounded concurrency and adaptive cooldown. - ``factory`` must be a zero-arg callable that returns a fresh coroutine - each call — a pre-created coroutine can only be awaited once, so on retry - the original would be dead. + The timeout covers concurrency queueing, 429 cooldown, backoff, and network + work together. ``factory`` must return a fresh coroutine for every retry. """ - limiter = _MintRateLimiter.get(mint_url) if mint_url else None + guard = _MintRateGuard.get(mint_url) if mint_url else None timeout = settings.mint_operation_timeout_seconds max_attempts = settings.mint_retry_max_attempts + 1 - last_exc: Exception | None = None - for attempt in range(max_attempts): - if limiter is not None: - await limiter.acquire() + async def invoke() -> Any: + if guard is not None: + return await guard.run(factory) + return await factory() - try: - if timeout > 0: - return await asyncio.wait_for(factory(), timeout=timeout) - return await factory() - except (asyncio.TimeoutError, httpx.TimeoutException) as exc: - last_exc = exc - if retry_timeouts and attempt < max_attempts - 1: - backoff = (2**attempt) + (time.monotonic() % 1.0) - logger.warning( - "Mint operation timed out, retrying", - extra={ - "op_name": op_name, - "mint_url": mint_url, - "attempt": attempt + 1, - "backoff_seconds": round(backoff, 2), - }, - ) - await asyncio.sleep(backoff) - continue - raise httpx.TimeoutException( - f"{op_name} timed out after {timeout}s (retried {attempt + 1}x)" - ) from exc - except httpx.HTTPStatusError as exc: - if _is_mint_rate_limited(exc) and attempt < max_attempts - 1: - backoff = (2**attempt) + (time.monotonic() % 1.0) - retry_after = _parse_retry_after(exc.response.headers) - if retry_after is not None: - backoff = max(retry_after, backoff) - logger.warning( - "Mint returned 429, backing off", - extra={ - "op_name": op_name, - "mint_url": mint_url, - "attempt": attempt + 1, - "backoff_seconds": round(backoff, 2), - }, - ) - await asyncio.sleep(backoff) - continue - raise - except Exception as exc: - if _is_mint_rate_limited(exc) and attempt < max_attempts - 1: - backoff = (2**attempt) + (time.monotonic() % 1.0) - logger.warning( - "Mint rate-limited, backing off", - extra={ - "op_name": op_name, - "mint_url": mint_url, - "attempt": attempt + 1, - "backoff_seconds": round(backoff, 2), - }, - ) - await asyncio.sleep(backoff) - continue - raise + async def run_with_retries() -> Any: + for attempt in range(max_attempts): + try: + return await invoke() + except (asyncio.TimeoutError, httpx.TimeoutException) as exc: + if retry_timeouts and attempt < max_attempts - 1: + backoff = (2**attempt) + (time.monotonic() % 1.0) + logger.warning( + "Mint operation timed out, retrying", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "backoff_seconds": round(backoff, 2), + }, + ) + await asyncio.sleep(backoff) + continue + raise httpx.TimeoutException( + f"{op_name} timed out (attempts: {attempt + 1})" + ) from exc + except Exception as exc: + if not _is_mint_rate_limited(exc): + raise - if last_exc: - raise last_exc - raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") + backoff = (2**attempt) + (time.monotonic() % 1.0) + if isinstance(exc, httpx.HTTPStatusError): + retry_after = _parse_retry_after(exc.response.headers) + if retry_after is not None: + backoff = max(retry_after, backoff) + if guard is not None: + guard.apply_cooldown(backoff) + + if attempt >= max_attempts - 1: + raise + logger.warning( + "Mint rate-limited, applying cooldown", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "cooldown_seconds": round(backoff, 2), + }, + ) + if guard is None: + await asyncio.sleep(backoff) + + raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") + + try: + if timeout > 0: + return await asyncio.wait_for(run_with_retries(), timeout=timeout) + return await run_with_retries() + except asyncio.TimeoutError as exc: + raise httpx.TimeoutException( + f"{op_name} exceeded its {timeout}s total timeout" + ) from exc def _parse_retry_after(headers: Any) -> float | None: diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 264bf26b..f712b357 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -27,15 +27,15 @@ def isolate_wallet_runtime_state() -> Generator[None, None, None]: from routstr import wallet as wallet_module from routstr.core.settings import settings - original_rpm = settings.mint_max_requests_per_minute - settings.mint_max_requests_per_minute = 0 - wallet_module._MintRateLimiter._limiters.clear() + original_concurrency = settings.mint_max_concurrency + settings.mint_max_concurrency = 0 + wallet_module._MintRateGuard._guards.clear() wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() wallet_module._wallet_load_locks.clear() yield - settings.mint_max_requests_per_minute = original_rpm - wallet_module._MintRateLimiter._limiters.clear() + settings.mint_max_concurrency = original_concurrency + wallet_module._MintRateGuard._guards.clear() wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() wallet_module._wallet_load_locks.clear() @@ -1342,7 +1342,7 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: # --------------------------------------------------------------------------- -# Per-mint limiter + _mint_operation factory/retry +# Per-mint adaptive guard + _mint_operation factory/retry # --------------------------------------------------------------------------- @@ -1366,42 +1366,54 @@ async def test_balance_proof_check_uses_large_batches_to_avoid_rate_limit() -> N @pytest.mark.asyncio -async def test_mint_rate_limiter_serializes_waiters_after_refill() -> None: - from routstr.wallet import _MintRateLimiter +async def test_mint_rate_guard_bounds_concurrency() -> None: + from routstr.wallet import _MintRateGuard - limiter = _MintRateLimiter("http://mint:3338", 60) - limiter._tokens = 0 - limiter._last_refill = 0 - clock = {"now": 0.0} - sleeps: list[float] = [] - real_sleep = asyncio.sleep + guard = _MintRateGuard("http://mint:3338", 2) + active = 0 + peak = 0 - async def fake_sleep(delay: float) -> None: - sleeps.append(delay) - clock["now"] += delay - await real_sleep(0) + async def operation() -> None: + nonlocal active, peak + active += 1 + peak = max(peak, active) + await asyncio.sleep(0) + active -= 1 - with patch("routstr.wallet.time.monotonic", side_effect=lambda: clock["now"]): - with patch("routstr.wallet.asyncio.sleep", side_effect=fake_sleep): - await asyncio.gather(limiter.acquire(), limiter.acquire()) + await asyncio.gather(*(guard.run(operation) for _ in range(5))) - assert sleeps == pytest.approx([1.0, 1.0]) - assert clock["now"] == pytest.approx(2.0) + assert peak == 2 -def test_mint_rate_limiter_rebuilds_when_setting_changes() -> None: +@pytest.mark.asyncio +async def test_mint_rate_guard_waits_for_adaptive_cooldown() -> None: + from routstr.wallet import _MintRateGuard + + guard = _MintRateGuard("http://mint:3338", 2) + guard._cooldown_until = 15.0 + operation = AsyncMock(return_value="ok") + + with patch("routstr.wallet.time.monotonic", return_value=10.0): + with patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep: + assert await guard.run(operation) == "ok" + + sleep.assert_awaited_once_with(5.0) + operation.assert_awaited_once() + + +def test_mint_rate_guard_rebuilds_when_setting_changes() -> None: from routstr.core.settings import settings - from routstr.wallet import _MintRateLimiter + from routstr.wallet import _MintRateGuard - with patch.object(settings, "mint_max_requests_per_minute", 20): - first = _MintRateLimiter.get("http://mint:3338") - with patch.object(settings, "mint_max_requests_per_minute", 10): - second = _MintRateLimiter.get("http://mint:3338") + with patch.object(settings, "mint_max_concurrency", 4): + first = _MintRateGuard.get("http://mint:3338") + with patch.object(settings, "mint_max_concurrency", 2): + second = _MintRateGuard.get("http://mint:3338") assert first is not None assert second is not None assert first is not second - assert second._max == 10 + assert second._max_concurrency == 2 @pytest.mark.asyncio @@ -1425,14 +1437,34 @@ async def test_mint_operation_honors_retry_after_as_minimum() -> None: sleep = AsyncMock() with patch.object(settings, "mint_retry_max_attempts", 1): with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch("routstr.wallet.time.monotonic", return_value=0.1): - with patch("routstr.wallet.asyncio.sleep", sleep): - result = await _mint_operation(factory, mint_url="http://mint:3338") + with patch.object(settings, "mint_max_concurrency", 1): + with patch("routstr.wallet.time.monotonic", return_value=0.1): + with patch("routstr.wallet.asyncio.sleep", sleep): + result = await _mint_operation( + factory, mint_url="http://mint:3338" + ) assert result == "ok" sleep.assert_awaited_once_with(60.0) +@pytest.mark.asyncio +async def test_mint_operation_timeout_includes_adaptive_cooldown() -> None: + from routstr.core.settings import settings + from routstr.wallet import _mint_operation, _MintRateGuard + + operation = AsyncMock(return_value="unexpected") + with patch.object(settings, "mint_max_concurrency", 1): + guard = _MintRateGuard.get("http://mint:3338") + assert guard is not None + guard.apply_cooldown(60) + with patch.object(settings, "mint_operation_timeout_seconds", 0.01): + with pytest.raises(httpx.TimeoutException, match="total timeout"): + await _mint_operation(operation, mint_url="http://mint:3338") + + operation.assert_not_awaited() + + @pytest.mark.asyncio async def test_mint_operation_retries_httpx_timeout_only_when_safe() -> None: from routstr.core.settings import settings @@ -1492,7 +1524,7 @@ async def test_mint_operation_factory_retry_succeeds() -> None: with patch.object(settings, "mint_retry_max_attempts", 3): with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_max_concurrency", 0): with patch("asyncio.sleep", AsyncMock()): result = await _mint_operation( factory, op_name="test_retry", mint_url="http://mint:3338" @@ -1518,7 +1550,7 @@ async def test_mint_operation_factory_retry_exhausted() -> None: with patch.object(settings, "mint_retry_max_attempts", 2): with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_max_concurrency", 0): with patch("asyncio.sleep", AsyncMock()): with pytest.raises(httpx.TimeoutException): await _mint_operation( @@ -1559,7 +1591,7 @@ async def test_lightning_mint_fallback_for_topups() -> None: with patch.object(settings, "primary_mint", primary): with patch.object(settings, "cashu_mints", [primary, secondary]): - with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_max_concurrency", 0): with patch.object(settings, "mint_operation_timeout_seconds", 0): with patch("routstr.lightning.get_wallet", side_effect=mock_get): bolt11, quote_id, mint_url = await _request_mint_with_fallback( @@ -1624,7 +1656,7 @@ async def test_swap_falls_back_to_secondary_mint() -> None: with patch.object(settings, "primary_mint", primary): with patch.object(settings, "primary_mint_unit", "sat"): with patch.object(settings, "cashu_mints", [primary, secondary]): - with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_max_concurrency", 0): with patch.object(settings, "mint_operation_timeout_seconds", 0): with patch("asyncio.sleep", AsyncMock()): with patch( @@ -1674,7 +1706,7 @@ async def test_lightning_mint_fallback_on_429() -> None: with patch.object(settings, "primary_mint", primary): with patch.object(settings, "cashu_mints", [primary, secondary]): with patch.object(settings, "mint_retry_max_attempts", 0): - with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_max_concurrency", 0): with patch.object(settings, "mint_operation_timeout_seconds", 0): with patch( "routstr.lightning.get_wallet", side_effect=mock_get @@ -1713,7 +1745,7 @@ async def test_lightning_mint_fallback_all_fail() -> None: with patch.object(settings, "primary_mint", primary): with patch.object(settings, "cashu_mints", [primary, secondary]): with patch.object(settings, "mint_retry_max_attempts", 0): - with patch.object(settings, "mint_max_requests_per_minute", 0): + with patch.object(settings, "mint_max_concurrency", 0): with patch.object(settings, "mint_operation_timeout_seconds", 0): with patch( "routstr.lightning.get_wallet", side_effect=mock_get From 40bf976fbc3de1c823ad5396080fb947aff5510f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 12 Jul 2026 15:04:43 +0200 Subject: [PATCH 08/46] fix: recreate mint URL migration --- ...21c84cd5ad83_add_mint_url_to_lightning_invoices.py} | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) rename migrations/versions/{add_mint_url_to_lightning_invoices.py => 21c84cd5ad83_add_mint_url_to_lightning_invoices.py} (66%) diff --git a/migrations/versions/add_mint_url_to_lightning_invoices.py b/migrations/versions/21c84cd5ad83_add_mint_url_to_lightning_invoices.py similarity index 66% rename from migrations/versions/add_mint_url_to_lightning_invoices.py rename to migrations/versions/21c84cd5ad83_add_mint_url_to_lightning_invoices.py index d3eb71b9..b69c16a8 100644 --- a/migrations/versions/add_mint_url_to_lightning_invoices.py +++ b/migrations/versions/21c84cd5ad83_add_mint_url_to_lightning_invoices.py @@ -1,13 +1,15 @@ -"""add mint_url to lightning_invoices +"""add mint url to lightning invoices -Revision ID: add_mint_url_li +Revision ID: 21c84cd5ad83 Revises: c6d7e8f9a0b1 -Create Date: 2026-07-10 02:00:00.000000 +Create Date: 2026-07-12 15:04:01.675455 """ + import sqlalchemy as sa from alembic import op -revision = "add_mint_url_li" +# revision identifiers, used by Alembic. +revision = "21c84cd5ad83" down_revision = "c6d7e8f9a0b1" branch_labels = None depends_on = None From 40153d4c36eeadb2bcbb63fd726d999c61f1a076 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 12 Jul 2026 15:07:31 +0200 Subject: [PATCH 09/46] fix: report cashu transaction persistence --- routstr/core/db.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/routstr/core/db.py b/routstr/core/db.py index 6cd5b593..1374b1ff 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -290,7 +290,7 @@ async def store_cashu_transaction( created_at: int | None = None, source: str = "x-cashu", api_key_hashed_key: str | None = None, -) -> None: +) -> bool: try: async with create_session() as session: tx = CashuTransaction( @@ -307,11 +307,13 @@ async def store_cashu_transaction( ) session.add(tx) await session.commit() + return True except Exception as e: logger.warning( f"Failed to store cashu transaction: {e} (type={typ})", extra={"error": str(e), "type": typ}, ) + return False class UpstreamProviderRow(SQLModel, table=True): # type: ignore From 65abcbce9258e716d634f3d68f30dbc0e6eeb251 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 13 Jul 2026 23:33:56 +0200 Subject: [PATCH 10/46] fix: harden mint fallback and refund recovery --- routstr/balance.py | 39 +++-- routstr/lightning.py | 27 +++- routstr/upstream/auto_topup.py | 25 +++- routstr/upstream/base.py | 7 +- routstr/wallet.py | 253 +++++++++++++++++++++++++------ tests/unit/test_auto_topup.py | 12 +- tests/unit/test_wallet.py | 264 ++++++++++++++++++++++++++++++++- 7 files changed, 561 insertions(+), 66 deletions(-) diff --git a/routstr/balance.py b/routstr/balance.py index cc37d089..8332577d 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -27,6 +27,7 @@ from .wallet import ( recieve_token, send_to_lnurl, send_token, + token_mint_url, ) router = APIRouter() @@ -220,7 +221,11 @@ async def _lookup_key_no_create( async def _restore_balance( - session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str + session: AsyncSession, + hashed_key: str, + balance: int, + reserved_balance: int, + mint_url: str, ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -235,7 +240,11 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, + extra={ + "hashed_key": hashed_key, + "restored_balance": balance, + "mint_url": mint_url, + }, ) @@ -389,15 +398,14 @@ async def refund_wallet_endpoint( detail="Balance changed concurrently. Please retry the refund.", ) - # --- MINT: balance is locked at zero, safe to create the refund token --- - # Proofs from untrusted mints are swapped to primary_mint on receive. - # Use primary_mint unless key.refund_mint_url is an explicitly trusted mint. + # The balance is locked at zero, so it is safe to create the refund token. effective_refund_mint = ( key.refund_mint_url if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints else settings.primary_mint ) try: + refund_currency = key.refund_currency or "sat" if key.refund_address: await send_to_lnurl( remaining_balance, @@ -407,10 +415,10 @@ async def refund_wallet_endpoint( ) result = {"recipient": key.refund_address} else: - refund_currency = key.refund_currency or "sat" token = await send_token( remaining_balance, refund_currency, effective_refund_mint ) + effective_refund_mint = token_mint_url(token, effective_refund_mint) result = {"token": token} if key.refund_currency == "sat": @@ -431,11 +439,23 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) error_msg = str(e) logger.error( "refund_wallet_endpoint: mint/send failed", @@ -462,7 +482,7 @@ async def refund_wallet_endpoint( token=result["token"], amount=remaining_balance, unit=key.refund_currency or "sat", - mint_url=key.refund_mint_url, + mint_url=effective_refund_mint, typ="out", collected=False, source="apikey", @@ -656,7 +676,6 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} - @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/lightning.py b/routstr/lightning.py index c7f0919b..fafffef9 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -14,6 +14,7 @@ from .core.settings import settings from .wallet import ( MintConnectionError, _is_mint_rate_limited, + _mint_cooldown_remaining, _mint_operation, get_wallet, is_mint_connection_error, @@ -75,17 +76,39 @@ async def _request_mint_with_fallback( *, allowed_mints: list[str] | None = None, ) -> tuple[str, str, str]: - """Request a quote, falling back only among the allowed trusted mints.""" + """Request a quote, falling back only among the allowed trusted mints. + + Guards against amount_sats <= 0: the cashu library's PostMintQuoteRequest + enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a + cryptic validation error deep in the stack. Fail fast with context. + """ + if amount_sats <= 0: + raise ValueError( + f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}." + ) tried: list[str] = [] configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints] candidates = list(dict.fromkeys(configured)) for mint_url in candidates: + cooldown = _mint_cooldown_remaining(mint_url) + if cooldown > 0: + tried.append(f"{mint_url}: cooling down") + logger.info( + "Skipping rate-limited mint", + extra={ + "mint_url": mint_url, + "cooldown_seconds": round(cooldown, 2), + "op_name": "request_mint_invoice", + }, + ) + continue try: - wallet = await get_wallet(mint_url, "sat") + wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False) quote = await _mint_operation( lambda: wallet.request_mint(amount_sats), op_name="request_mint_invoice", mint_url=mint_url, + retry_on_rate_limit=False, ) return quote.request, quote.quote, mint_url except Exception as e: diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index 3517be7d..932e3263 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -10,7 +10,7 @@ from ..core.db import ( create_session, store_cashu_transaction, ) -from ..wallet import send_token +from ..wallet import release_token_reservation, send_token, token_mint_url from .routstr import RoutstrUpstreamProvider logger = get_logger(__name__) @@ -142,20 +142,33 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) return + actual_mint_url = token_mint_url(token, mint_url) stored = await store_cashu_transaction( token=token, amount=amount, unit="sat", - mint_url=mint_url, + mint_url=actual_mint_url, typ="out", collected=False, source="auto_topup", ) if not stored: - logger.critical( - "Aborting auto top-up because its cashu token could not be persisted", - extra={"provider_id": row.id, "mint_url": mint_url}, - ) + try: + await release_token_reservation(token) + except Exception as error: + logger.critical( + "Failed to release untracked auto-topup token", + extra={ + "provider_id": row.id, + "mint_url": actual_mint_url, + "error": str(error), + }, + ) + else: + logger.warning( + "Auto-topup token was released after persistence failed", + extra={"provider_id": row.id, "mint_url": actual_mint_url}, + ) return result = await provider.topup(token) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index f8ea2d4a..bbbe0f49 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -45,6 +45,7 @@ from ..wallet import ( classify_redemption_error, recieve_token, send_token, + token_mint_url, ) from . import messages_dispatch from .cache_breakpoints import ( @@ -3290,7 +3291,7 @@ class BaseUpstreamProvider: token=refund_token, amount=amount, unit=unit, - mint_url=mint, + mint_url=token_mint_url(refund_token, mint), typ="out", request_id=request_id, ) @@ -3645,7 +3646,7 @@ class BaseUpstreamProvider: token=refund_token, amount=emergency_refund, unit=unit, - mint_url=mint, + mint_url=token_mint_url(refund_token, mint), typ="out", request_id=request_id, ) @@ -4609,7 +4610,7 @@ class BaseUpstreamProvider: token=refund_token, amount=emergency_refund, unit=unit, - mint_url=mint, + mint_url=token_mint_url(refund_token, mint), typ="out", request_id=request_id, ) diff --git a/routstr/wallet.py b/routstr/wallet.py index 273d1aaf..f992c7a0 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -63,15 +63,13 @@ _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( class _MintRateGuard: - """Bound concurrency and adapt to actual per-mint 429 responses.""" + """Limit concurrency and remember per-mint rate-limit cooldowns.""" _guards: dict[str, "_MintRateGuard"] = {} @classmethod - def get(cls, mint_url: str) -> "_MintRateGuard | None": + def get(cls, mint_url: str) -> "_MintRateGuard": concurrency = settings.mint_max_concurrency - if concurrency <= 0: - return None guard = cls._guards.get(mint_url) if guard is None or guard._max_concurrency != concurrency: guard = cls(mint_url, concurrency) @@ -81,7 +79,9 @@ class _MintRateGuard: def __init__(self, mint_url: str, max_concurrency: int): self._mint_url = mint_url self._max_concurrency = max_concurrency - self._semaphore = asyncio.Semaphore(max_concurrency) + self._semaphore = ( + asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None + ) self._cooldown_until = 0.0 def apply_cooldown(self, delay: float) -> None: @@ -89,19 +89,28 @@ class _MintRateGuard: self._cooldown_until, time.monotonic() + max(0.0, delay) ) + def cooldown_remaining(self) -> float: + return max(0.0, self._cooldown_until - time.monotonic()) + + async def _run_after_cooldown(self, factory: Callable[[], Awaitable[Any]]) -> Any: + wait = self.cooldown_remaining() + if wait > 0: + logger.debug( + "Mint rate guard: cooling down", + extra={"mint_url": self._mint_url, "wait_seconds": round(wait, 2)}, + ) + await asyncio.sleep(wait) + return await factory() + async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any: + if self._semaphore is None: + return await self._run_after_cooldown(factory) async with self._semaphore: - wait = self._cooldown_until - time.monotonic() - if wait > 0: - logger.debug( - "Mint rate guard: cooling down", - extra={ - "mint_url": self._mint_url, - "wait_seconds": round(wait, 2), - }, - ) - await asyncio.sleep(wait) - return await factory() + return await self._run_after_cooldown(factory) + + +def _mint_cooldown_remaining(mint_url: str) -> float: + return _MintRateGuard.get(mint_url).cooldown_remaining() def _is_mint_rate_limited(error: BaseException) -> bool: @@ -126,11 +135,17 @@ async def _mint_operation( op_name: str = "mint_operation", mint_url: str = "", retry_timeouts: bool = True, + retry_on_rate_limit: bool = True, ) -> Any: """Run a mint operation with bounded concurrency and adaptive cooldown. The timeout covers concurrency queueing, 429 cooldown, backoff, and network - work together. ``factory`` must return a fresh coroutine for every retry. + work together. ``factory`` must return a fresh coroutine for every retry. + + When ``retry_on_rate_limit`` is False a 429 is not retried in-place — the + cooldown is still applied to the per-mint guard (so subsequent operations on + that mint wait), but the exception is re-raised so the caller (typically + ``_request_mint_with_fallback``) can immediately try a different mint. """ guard = _MintRateGuard.get(mint_url) if mint_url else None timeout = settings.mint_operation_timeout_seconds @@ -166,6 +181,9 @@ async def _mint_operation( if not _is_mint_rate_limited(exc): raise + # Apply cooldown to the guard regardless — even when we're + # about to re-raise for fallback, the guard must remember that + # this mint is rate-limited for future operations. backoff = (2**attempt) + (time.monotonic() % 1.0) if isinstance(exc, httpx.HTTPStatusError): retry_after = _parse_retry_after(exc.response.headers) @@ -174,6 +192,20 @@ async def _mint_operation( if guard is not None: guard.apply_cooldown(backoff) + # When the caller has a fallback strategy (trusted-mint + # list), re-raise immediately so the caller can try the next + # mint instead of waiting through this mint's cooldown. + if not retry_on_rate_limit: + logger.warning( + "Mint rate-limited, skipping retries for fallback", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "cooldown_seconds": round(backoff, 2), + }, + ) + raise + if attempt >= max_attempts - 1: raise logger.warning( @@ -374,24 +406,13 @@ async def recieve_token( async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: - """Internal send function - returns amount and serialized token""" - effective_mint_url = mint_url or settings.primary_mint - wallet: Wallet = await get_wallet(effective_mint_url, unit) - proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) - proofs_for_mint = sum(p.amount for p in proofs) - - # Fallback: proofs from untrusted source mints are swapped to primary_mint - # during receive, so the user's preferred refund_mint_url may have no proofs - # even though the global wallet has the balance. - if proofs_for_mint < amount and effective_mint_url != settings.primary_mint: - logger.info( - f"send: insufficient proofs at {effective_mint_url} " - f"(have {proofs_for_mint}, need {amount}), falling back to primary_mint={settings.primary_mint}" - ) - effective_mint_url = settings.primary_mint - wallet = await get_wallet(effective_mint_url, unit) - proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) - proofs_for_mint = sum(p.amount for p in proofs) + """Create a token from the preferred mint or another funded trusted mint.""" + effective_mint_url = await find_trusted_mint_with_funds(amount, unit, mint_url) + wallet = await get_wallet(effective_mint_url, unit) + proofs = get_proofs_per_mint_and_unit( + wallet, effective_mint_url, unit, not_reserved=True + ) + proofs_for_mint = sum(proof.amount for proof in proofs) all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { @@ -435,6 +456,61 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str return token +async def release_token_reservation(token: str) -> None: + """Release a token that was created locally but never handed off.""" + token_obj = deserialize_token_from_string(token) + wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) + await wallet.set_reserved_for_send(token_obj.proofs, reserved=False) + + secrets = {proof.secret for proof in token_obj.proofs} + for proof in token_obj.proofs: + proof.reserved = False + for proof in wallet.proofs: + if proof.secret in secrets: + proof.reserved = False + + +def token_mint_url(token: str, fallback: str | None = None) -> str: + try: + return str(deserialize_token_from_string(token).mint) + except Exception: + if fallback is None: + raise + return fallback + + +async def find_trusted_mint_with_funds( + amount: int, unit: str, preferred_mint: str | None = None +) -> str: + """Choose a trusted mint that can cover a refund without waiting on cooldown.""" + trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) + candidates: list[str] = [] + if preferred_mint in trusted: + candidates.append(preferred_mint) + candidates.extend(mint for mint in trusted if mint not in candidates) + + balances: dict[str, int] = {} + for mint_url in candidates: + if _mint_cooldown_remaining(mint_url) > 0: + continue + try: + wallet = await get_wallet(mint_url, unit, retry_on_rate_limit=False) + except Exception as error: + if is_mint_connection_error(error) or _is_mint_rate_limited(error): + balances[mint_url] = 0 + continue + raise + + proofs = get_proofs_per_mint_and_unit(wallet, mint_url, unit, not_reserved=True) + balances[mint_url] = sum(proof.amount for proof in proofs) + if balances[mint_url] >= amount: + return mint_url + + raise ValueError( + f"No trusted mint has {amount} {unit} available; balances={balances}" + ) + + # A foreign mint's fee_reserve is a non-binding estimate (NUT-05): the mint may # demand more when re-quoting or at melt execution. Instead of padding the # estimate with a safety buffer (which strands the margin at the foreign mint @@ -505,21 +581,48 @@ async def _request_mint_with_fallback( amount: int, *, op_name: str, primary_wallet: Wallet | None = None ) -> tuple[Wallet, str, MintQuote]: """Try request_mint on the primary mint, fall back to other trusted mints - on transport or rate-limit failure. Returns the wallet, mint_url, and quote.""" + on transport or rate-limit failure. Returns the wallet, mint_url, and quote. + + Guards against amount <= 0: the cashu library's PostMintQuoteRequest + enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a + cryptic validation error deep in the stack. Fail fast with context. + """ + if amount <= 0: + raise ValueError( + f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " + f"Token value is too small after fee deduction or unit conversion." + ) candidates = [settings.primary_mint] + [ m for m in settings.cashu_mints if m != settings.primary_mint ] tried: list[str] = [] for mint_url in candidates: + cooldown = _mint_cooldown_remaining(mint_url) + if cooldown > 0: + tried.append(f"{mint_url}: cooling down") + logger.info( + "Skipping rate-limited mint", + extra={ + "mint_url": mint_url, + "cooldown_seconds": round(cooldown, 2), + "op_name": op_name, + }, + ) + continue try: if mint_url == settings.primary_mint and primary_wallet is not None: wallet = primary_wallet else: - wallet = await get_wallet(mint_url, settings.primary_mint_unit) + wallet = await get_wallet( + mint_url, + settings.primary_mint_unit, + retry_on_rate_limit=False, + ) quote = await _mint_operation( lambda: wallet.request_mint(amount), op_name=op_name, mint_url=mint_url, + retry_on_rate_limit=False, ) return wallet, mint_url, quote except Exception as e: @@ -563,11 +666,37 @@ async def _calculate_swap_amount( ) return int(receive_amount) + # The cashu library's PostMintQuoteRequest enforces amount > 0 (Pydantic + # Field(gt=0)). When the token's face value in the primary mint's unit + # truncates to 0 (e.g. < 1000 msat with a "sat" primary unit), calling + # request_mint(0) raises a validation error that is cryptic in production + # logs. Guard early with full diagnostic context instead. + if receive_amount <= 0: + logger.error( + "swap_to_primary_mint: receive_amount is zero or negative, cannot estimate fees", + extra={ + "amount_msat": amount_msat, + "token_unit": token_unit, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, + "primary_mint_unit": settings.primary_mint_unit, + "receive_amount": receive_amount, + }, + ) + raise ValueError( + f"Token amount ({amount_msat} msat, unit={token_unit}) is too small to " + f"swap to primary mint ({settings.primary_mint}, unit={settings.primary_mint_unit}): " + f"receive_amount={receive_amount}. Minimum 1 {settings.primary_mint_unit} required." + ) + logger.info( "swap_to_primary_mint: estimating fees", extra={ "dummy_amount": receive_amount, "unit": settings.primary_mint_unit, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, + "amount_msat": amount_msat, }, ) @@ -600,6 +729,9 @@ async def _calculate_swap_amount( "input_fees": input_fees, "minted_amount": minted_amount, "minted_unit": settings.primary_mint_unit, + "fee_reserve": fee_reserve, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, }, ) return minted_amount @@ -607,7 +739,16 @@ async def _calculate_swap_amount( except Exception as e: logger.error( "swap_to_primary_mint: fee estimation failed", - extra={"error": str(e)}, + extra={ + "error": str(e), + "error_type": type(e).__name__, + "amount_msat": amount_msat, + "token_unit": token_unit, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, + "primary_mint_unit": settings.primary_mint_unit, + "receive_amount": receive_amount, + }, ) if is_mint_connection_error(e): raise MintConnectionError("Cashu mint is unreachable") from e @@ -672,6 +813,24 @@ async def swap_to_primary_mint( dest_mint_url = settings.primary_mint while True: attempt += 1 + if minted_amount <= 0: + logger.error( + "swap_to_primary_mint: minted_amount is zero or negative before requesting quote", + extra={ + "minted_amount": minted_amount, + "attempt": attempt, + "foreign_mint": token_obj.mint, + "token_amount": token_amount, + "token_unit": token_obj.unit, + "amount_msat": amount_msat, + "observed_extra_fee": observed_extra_fee, + "primary_mint": settings.primary_mint, + }, + ) + raise ValueError( + f"Cannot swap token ({token_amount} {token_obj.unit}) from {token_obj.mint}: " + f"minted_amount={minted_amount} after fee deduction (attempt {attempt})" + ) dest_wallet, dest_mint_url, mint_quote = await _request_mint_with_fallback( minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet ) @@ -1006,7 +1165,12 @@ _wallet_load_locks: dict[str, asyncio.Lock] = {} _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30 -async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet: +async def get_wallet( + mint_url: str, + unit: str = "sat", + load: bool = True, + retry_on_rate_limit: bool = True, +) -> Wallet: global _wallets, _wallet_last_load, _wallet_load_locks id = f"{mint_url}_{unit}" lock = _wallet_load_locks.setdefault(id, asyncio.Lock()) @@ -1016,17 +1180,19 @@ async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wal if load: now = time.monotonic() - last = _wallet_last_load.get(id, 0) - if now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: + last = _wallet_last_load.get(id) + if last is None or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: await _mint_operation( lambda: _wallets[id].load_mint(), op_name="load_mint", mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, ) await _mint_operation( lambda: _wallets[id].load_proofs(reload=True), op_name="load_proofs", mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, ) _wallet_last_load[id] = time.monotonic() return _wallets[id] @@ -1361,9 +1527,10 @@ async def periodic_routstr_fee_payout() -> None: async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: + mint = await find_trusted_mint_with_funds(amount, unit, mint) wallet = await get_wallet(mint, unit) - proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] - proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) + available = get_proofs_per_mint_and_unit(wallet, mint, unit, not_reserved=True) + proofs, _ = await wallet.select_to_send(available, amount, set_reserved=True) return await raw_send_to_lnurl(wallet, proofs, address, unit) diff --git a/tests/unit/test_auto_topup.py b/tests/unit/test_auto_topup.py index c05c85e5..cf3d06e2 100644 --- a/tests/unit/test_auto_topup.py +++ b/tests/unit/test_auto_topup.py @@ -66,6 +66,10 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected() "routstr.upstream.auto_topup.store_cashu_transaction", AsyncMock(return_value=True), ) as store, + patch( + "routstr.upstream.auto_topup.token_mint_url", + return_value="https://fallback-mint.test", + ), patch("routstr.upstream.auto_topup.create_session", return_value=session), ): await _check_and_topup(_row()) @@ -74,7 +78,7 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected() token="cashu-token", amount=50, unit="sat", - mint_url="https://mint.test", + mint_url="https://fallback-mint.test", typ="out", collected=False, source="auto_topup", @@ -138,6 +142,12 @@ async def test_auto_topup_does_not_send_untracked_token() -> None: "routstr.upstream.auto_topup.store_cashu_transaction", AsyncMock(return_value=False), ), + patch( + "routstr.upstream.auto_topup.release_token_reservation", + AsyncMock(), + ) as reclaim, ): await _check_and_topup(_row()) + + reclaim.assert_awaited_once_with("cashu-token") provider.topup.assert_not_awaited() diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index f712b357..814fb36c 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -169,6 +169,59 @@ async def test_send_token() -> None: assert token == "test_token" +@pytest.mark.asyncio +async def test_release_token_reservation_unreserves_local_proofs() -> None: + from routstr.wallet import release_token_reservation + + token_proof = Mock(secret="proof-secret", reserved=True) + cached_proof = Mock(secret="proof-secret", reserved=True) + token = Mock(mint="http://mint:3338", unit="sat", proofs=[token_proof]) + wallet = Mock(proofs=[cached_proof], set_reserved_for_send=AsyncMock()) + with ( + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch( + "routstr.wallet.get_wallet", AsyncMock(return_value=wallet) + ) as get_wallet, + ): + await release_token_reservation("cashu-token") + + get_wallet.assert_awaited_once_with("http://mint:3338", "sat", load=False) + wallet.set_reserved_for_send.assert_awaited_once_with(token.proofs, reserved=False) + assert token_proof.reserved is False + assert cached_proof.reserved is False + + +@pytest.mark.asyncio +async def test_refund_mint_falls_back_to_trusted_mint_with_funds() -> None: + from routstr.core.settings import settings + from routstr.wallet import find_trusted_mint_with_funds + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + def wallet_for(mint: str, amount: int) -> Mock: + keyset = Mock(id=f"keyset-{mint}", mint_url=mint) + keyset.unit.name = "sat" + proof = Mock(id=keyset.id, amount=amount, reserved=False) + return Mock(keysets={keyset.id: keyset}, proofs=[proof]) + + wallets = { + primary: wallet_for(primary, 50), + secondary: wallet_for(secondary, 200), + } + with ( + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary, secondary]), + patch( + "routstr.wallet.get_wallet", + AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), + ), + ): + mint = await find_trusted_mint_with_funds(100, "sat", primary) + + assert mint == secondary + + @pytest.mark.asyncio async def test_credit_balance() -> None: token_data = { @@ -1416,6 +1469,25 @@ def test_mint_rate_guard_rebuilds_when_setting_changes() -> None: assert second._max_concurrency == 2 +@pytest.mark.asyncio +async def test_mint_rate_guard_keeps_cooldown_when_concurrency_is_unlimited() -> None: + from routstr.core.settings import settings + from routstr.wallet import _MintRateGuard + + operation = AsyncMock(return_value="ok") + with ( + patch.object(settings, "mint_max_concurrency", 0), + patch("routstr.wallet.time.monotonic", return_value=0), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + ): + guard = _MintRateGuard.get("http://mint:3338") + guard.apply_cooldown(5) + assert await guard.run(operation) == "ok" + + sleep.assert_awaited_once_with(5) + operation.assert_awaited_once() + + @pytest.mark.asyncio async def test_mint_operation_honors_retry_after_as_minimum() -> None: from routstr.core.settings import settings @@ -1495,7 +1567,9 @@ async def test_get_wallet_initializes_and_loads_once_concurrently() -> None: with patch( "routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet) ) as create: - with patch("routstr.wallet.time.monotonic", return_value=100.0): + # A fresh wallet must load even when the host has been up for less than + # the reload interval. + with patch("routstr.wallet.time.monotonic", return_value=10.0): first, second = await asyncio.gather( get_wallet("http://mint:3338"), get_wallet("http://mint:3338") ) @@ -1506,6 +1580,36 @@ async def test_get_wallet_initializes_and_loads_once_concurrently() -> None: mock_wallet.load_proofs.assert_awaited_once_with(reload=True) +@pytest.mark.asyncio +async def test_get_wallet_can_surface_429_without_retrying() -> None: + from routstr.core.settings import settings + from routstr.wallet import get_wallet + + request = httpx.Request("GET", "http://mint:3338/v1/info") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + wallet = Mock( + load_mint=AsyncMock( + side_effect=httpx.HTTPStatusError( + "rate limited", request=request, response=response + ) + ), + load_proofs=AsyncMock(), + ) + + with ( + patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=wallet)), + patch.object(settings, "mint_retry_max_attempts", 3), + patch.object(settings, "mint_operation_timeout_seconds", 0), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + ): + with pytest.raises(httpx.HTTPStatusError): + await get_wallet("http://mint:3338", retry_on_rate_limit=False) + + wallet.load_mint.assert_awaited_once() + wallet.load_proofs.assert_not_awaited() + sleep.assert_not_awaited() + + @pytest.mark.asyncio async def test_mint_operation_factory_retry_succeeds() -> None: """_mint_operation accepts a zero-arg factory, not a dead coroutine. @@ -1752,3 +1856,161 @@ async def test_lightning_mint_fallback_all_fail() -> None: ): with pytest.raises(MintConnectionError): await _request_mint_with_fallback(1000) + + +@pytest.mark.asyncio +async def test_lightning_mint_fallback_rejects_zero_amount() -> None: + """Zero or negative amounts must be rejected before reaching the mint.""" + from routstr.lightning import _request_mint_with_fallback + + with pytest.raises(ValueError, match="amount_sats must be > 0"): + await _request_mint_with_fallback(0) + + with pytest.raises(ValueError, match="amount_sats must be > 0"): + await _request_mint_with_fallback(-5) + + +@pytest.mark.asyncio +async def test_wallet_request_mint_fallback_rejects_zero_amount() -> None: + """Zero or negative amounts must be rejected before reaching the mint.""" + from routstr.wallet import _request_mint_with_fallback + + with pytest.raises(ValueError, match="amount must be > 0"): + await _request_mint_with_fallback(0, op_name="test") + + with pytest.raises(ValueError, match="amount must be > 0"): + await _request_mint_with_fallback(-1, op_name="test") + + +@pytest.mark.asyncio +async def test_wallet_fallback_on_429_no_in_place_retry() -> None: + """A 429 from the primary mint must trigger immediate fallback to the + secondary — _mint_operation must NOT retry in-place when + retry_on_rate_limit=False is set by _request_mint_with_fallback.""" + from routstr.core.settings import settings + from routstr.wallet import _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + request = httpx.Request("POST", "http://primary:3338/v1/mint/quote/bolt11") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + primary_call_count = 0 + + async def primary_request_mint(amount): + nonlocal primary_call_count + primary_call_count += 1 + raise httpx.HTTPStatusError("rate limited", request=request, response=response) + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock(side_effect=primary_request_mint) + + mock_quote = Mock(quote="q_secondary", request="lnbc1secondary") + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_retry_max_attempts", 3): + with patch.object(settings, "mint_max_concurrency", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("asyncio.sleep", AsyncMock()) as mock_sleep: + with patch( + "routstr.wallet.get_wallet", side_effect=mock_get + ): + _, mint_url, _ = await _request_mint_with_fallback( + 1000, op_name="test_429_fallback" + ) + + assert mint_url == secondary + assert primary_call_count == 1 + mock_secondary_wallet.request_mint.assert_called_once() + mock_sleep.assert_not_called() + + +@pytest.mark.asyncio +async def test_wallet_fallback_skips_mint_during_cooldown() -> None: + from routstr.core.settings import settings + from routstr.wallet import _MintRateGuard, _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + primary_wallet = Mock(request_mint=AsyncMock()) + quote = Mock(quote="q_secondary", request="lnbc1secondary") + secondary_wallet = Mock(request_mint=AsyncMock(return_value=quote)) + wallets = {primary: primary_wallet, secondary: secondary_wallet} + + with ( + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary, secondary]), + patch.object(settings, "mint_max_concurrency", 0), + patch.object(settings, "mint_operation_timeout_seconds", 0), + patch("routstr.wallet.time.monotonic", return_value=10), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + patch( + "routstr.wallet.get_wallet", + AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), + ), + ): + _MintRateGuard.get(primary).apply_cooldown(60) + _, mint_url, _ = await _request_mint_with_fallback( + 1000, op_name="test_cooldown_fallback" + ) + + assert mint_url == secondary + primary_wallet.request_mint.assert_not_awaited() + secondary_wallet.request_mint.assert_awaited_once_with(1000) + sleep.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_lightning_fallback_on_429_no_in_place_retry() -> None: + """Same as above but for the lightning.py _request_mint_with_fallback.""" + from routstr.core.settings import settings + from routstr.lightning import _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + request = httpx.Request("POST", "http://primary:3338/v1/mint/quote/bolt11") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + primary_call_count = 0 + + async def primary_request_mint(amount): + nonlocal primary_call_count + primary_call_count += 1 + raise httpx.HTTPStatusError("rate limited", request=request, response=response) + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock(side_effect=primary_request_mint) + + mock_quote = Mock(quote="q_secondary", request="lnbc1secondary") + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_retry_max_attempts", 3): + with patch.object(settings, "mint_max_concurrency", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("asyncio.sleep", AsyncMock()) as mock_sleep: + with patch( + "routstr.lightning.get_wallet", side_effect=mock_get + ): + _, _, first_mint = await _request_mint_with_fallback( + 1000 + ) + _, _, second_mint = await _request_mint_with_fallback( + 1000 + ) + + assert first_mint == second_mint == secondary + assert primary_call_count == 1 + assert mock_secondary_wallet.request_mint.await_count == 2 + mock_sleep.assert_not_called() From 65702171e404238a0ac324768ae511f4a69b7691 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 13 Jul 2026 23:33:56 +0200 Subject: [PATCH 11/46] fix: harden mint fallback and refund recovery --- routstr/balance.py | 39 +++-- routstr/lightning.py | 27 +++- routstr/upstream/auto_topup.py | 25 +++- routstr/upstream/base.py | 7 +- routstr/wallet.py | 253 +++++++++++++++++++++++++------ tests/unit/test_auto_topup.py | 12 +- tests/unit/test_wallet.py | 264 ++++++++++++++++++++++++++++++++- 7 files changed, 561 insertions(+), 66 deletions(-) diff --git a/routstr/balance.py b/routstr/balance.py index cc37d089..8332577d 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -27,6 +27,7 @@ from .wallet import ( recieve_token, send_to_lnurl, send_token, + token_mint_url, ) router = APIRouter() @@ -220,7 +221,11 @@ async def _lookup_key_no_create( async def _restore_balance( - session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str + session: AsyncSession, + hashed_key: str, + balance: int, + reserved_balance: int, + mint_url: str, ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -235,7 +240,11 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, + extra={ + "hashed_key": hashed_key, + "restored_balance": balance, + "mint_url": mint_url, + }, ) @@ -389,15 +398,14 @@ async def refund_wallet_endpoint( detail="Balance changed concurrently. Please retry the refund.", ) - # --- MINT: balance is locked at zero, safe to create the refund token --- - # Proofs from untrusted mints are swapped to primary_mint on receive. - # Use primary_mint unless key.refund_mint_url is an explicitly trusted mint. + # The balance is locked at zero, so it is safe to create the refund token. effective_refund_mint = ( key.refund_mint_url if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints else settings.primary_mint ) try: + refund_currency = key.refund_currency or "sat" if key.refund_address: await send_to_lnurl( remaining_balance, @@ -407,10 +415,10 @@ async def refund_wallet_endpoint( ) result = {"recipient": key.refund_address} else: - refund_currency = key.refund_currency or "sat" token = await send_token( remaining_balance, refund_currency, effective_refund_mint ) + effective_refund_mint = token_mint_url(token, effective_refund_mint) result = {"token": token} if key.refund_currency == "sat": @@ -431,11 +439,23 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) error_msg = str(e) logger.error( "refund_wallet_endpoint: mint/send failed", @@ -462,7 +482,7 @@ async def refund_wallet_endpoint( token=result["token"], amount=remaining_balance, unit=key.refund_currency or "sat", - mint_url=key.refund_mint_url, + mint_url=effective_refund_mint, typ="out", collected=False, source="apikey", @@ -656,7 +676,6 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} - @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/lightning.py b/routstr/lightning.py index c7f0919b..fafffef9 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -14,6 +14,7 @@ from .core.settings import settings from .wallet import ( MintConnectionError, _is_mint_rate_limited, + _mint_cooldown_remaining, _mint_operation, get_wallet, is_mint_connection_error, @@ -75,17 +76,39 @@ async def _request_mint_with_fallback( *, allowed_mints: list[str] | None = None, ) -> tuple[str, str, str]: - """Request a quote, falling back only among the allowed trusted mints.""" + """Request a quote, falling back only among the allowed trusted mints. + + Guards against amount_sats <= 0: the cashu library's PostMintQuoteRequest + enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a + cryptic validation error deep in the stack. Fail fast with context. + """ + if amount_sats <= 0: + raise ValueError( + f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}." + ) tried: list[str] = [] configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints] candidates = list(dict.fromkeys(configured)) for mint_url in candidates: + cooldown = _mint_cooldown_remaining(mint_url) + if cooldown > 0: + tried.append(f"{mint_url}: cooling down") + logger.info( + "Skipping rate-limited mint", + extra={ + "mint_url": mint_url, + "cooldown_seconds": round(cooldown, 2), + "op_name": "request_mint_invoice", + }, + ) + continue try: - wallet = await get_wallet(mint_url, "sat") + wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False) quote = await _mint_operation( lambda: wallet.request_mint(amount_sats), op_name="request_mint_invoice", mint_url=mint_url, + retry_on_rate_limit=False, ) return quote.request, quote.quote, mint_url except Exception as e: diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index 3517be7d..932e3263 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -10,7 +10,7 @@ from ..core.db import ( create_session, store_cashu_transaction, ) -from ..wallet import send_token +from ..wallet import release_token_reservation, send_token, token_mint_url from .routstr import RoutstrUpstreamProvider logger = get_logger(__name__) @@ -142,20 +142,33 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) return + actual_mint_url = token_mint_url(token, mint_url) stored = await store_cashu_transaction( token=token, amount=amount, unit="sat", - mint_url=mint_url, + mint_url=actual_mint_url, typ="out", collected=False, source="auto_topup", ) if not stored: - logger.critical( - "Aborting auto top-up because its cashu token could not be persisted", - extra={"provider_id": row.id, "mint_url": mint_url}, - ) + try: + await release_token_reservation(token) + except Exception as error: + logger.critical( + "Failed to release untracked auto-topup token", + extra={ + "provider_id": row.id, + "mint_url": actual_mint_url, + "error": str(error), + }, + ) + else: + logger.warning( + "Auto-topup token was released after persistence failed", + extra={"provider_id": row.id, "mint_url": actual_mint_url}, + ) return result = await provider.topup(token) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index f8ea2d4a..bbbe0f49 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -45,6 +45,7 @@ from ..wallet import ( classify_redemption_error, recieve_token, send_token, + token_mint_url, ) from . import messages_dispatch from .cache_breakpoints import ( @@ -3290,7 +3291,7 @@ class BaseUpstreamProvider: token=refund_token, amount=amount, unit=unit, - mint_url=mint, + mint_url=token_mint_url(refund_token, mint), typ="out", request_id=request_id, ) @@ -3645,7 +3646,7 @@ class BaseUpstreamProvider: token=refund_token, amount=emergency_refund, unit=unit, - mint_url=mint, + mint_url=token_mint_url(refund_token, mint), typ="out", request_id=request_id, ) @@ -4609,7 +4610,7 @@ class BaseUpstreamProvider: token=refund_token, amount=emergency_refund, unit=unit, - mint_url=mint, + mint_url=token_mint_url(refund_token, mint), typ="out", request_id=request_id, ) diff --git a/routstr/wallet.py b/routstr/wallet.py index 273d1aaf..f992c7a0 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -63,15 +63,13 @@ _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( class _MintRateGuard: - """Bound concurrency and adapt to actual per-mint 429 responses.""" + """Limit concurrency and remember per-mint rate-limit cooldowns.""" _guards: dict[str, "_MintRateGuard"] = {} @classmethod - def get(cls, mint_url: str) -> "_MintRateGuard | None": + def get(cls, mint_url: str) -> "_MintRateGuard": concurrency = settings.mint_max_concurrency - if concurrency <= 0: - return None guard = cls._guards.get(mint_url) if guard is None or guard._max_concurrency != concurrency: guard = cls(mint_url, concurrency) @@ -81,7 +79,9 @@ class _MintRateGuard: def __init__(self, mint_url: str, max_concurrency: int): self._mint_url = mint_url self._max_concurrency = max_concurrency - self._semaphore = asyncio.Semaphore(max_concurrency) + self._semaphore = ( + asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None + ) self._cooldown_until = 0.0 def apply_cooldown(self, delay: float) -> None: @@ -89,19 +89,28 @@ class _MintRateGuard: self._cooldown_until, time.monotonic() + max(0.0, delay) ) + def cooldown_remaining(self) -> float: + return max(0.0, self._cooldown_until - time.monotonic()) + + async def _run_after_cooldown(self, factory: Callable[[], Awaitable[Any]]) -> Any: + wait = self.cooldown_remaining() + if wait > 0: + logger.debug( + "Mint rate guard: cooling down", + extra={"mint_url": self._mint_url, "wait_seconds": round(wait, 2)}, + ) + await asyncio.sleep(wait) + return await factory() + async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any: + if self._semaphore is None: + return await self._run_after_cooldown(factory) async with self._semaphore: - wait = self._cooldown_until - time.monotonic() - if wait > 0: - logger.debug( - "Mint rate guard: cooling down", - extra={ - "mint_url": self._mint_url, - "wait_seconds": round(wait, 2), - }, - ) - await asyncio.sleep(wait) - return await factory() + return await self._run_after_cooldown(factory) + + +def _mint_cooldown_remaining(mint_url: str) -> float: + return _MintRateGuard.get(mint_url).cooldown_remaining() def _is_mint_rate_limited(error: BaseException) -> bool: @@ -126,11 +135,17 @@ async def _mint_operation( op_name: str = "mint_operation", mint_url: str = "", retry_timeouts: bool = True, + retry_on_rate_limit: bool = True, ) -> Any: """Run a mint operation with bounded concurrency and adaptive cooldown. The timeout covers concurrency queueing, 429 cooldown, backoff, and network - work together. ``factory`` must return a fresh coroutine for every retry. + work together. ``factory`` must return a fresh coroutine for every retry. + + When ``retry_on_rate_limit`` is False a 429 is not retried in-place — the + cooldown is still applied to the per-mint guard (so subsequent operations on + that mint wait), but the exception is re-raised so the caller (typically + ``_request_mint_with_fallback``) can immediately try a different mint. """ guard = _MintRateGuard.get(mint_url) if mint_url else None timeout = settings.mint_operation_timeout_seconds @@ -166,6 +181,9 @@ async def _mint_operation( if not _is_mint_rate_limited(exc): raise + # Apply cooldown to the guard regardless — even when we're + # about to re-raise for fallback, the guard must remember that + # this mint is rate-limited for future operations. backoff = (2**attempt) + (time.monotonic() % 1.0) if isinstance(exc, httpx.HTTPStatusError): retry_after = _parse_retry_after(exc.response.headers) @@ -174,6 +192,20 @@ async def _mint_operation( if guard is not None: guard.apply_cooldown(backoff) + # When the caller has a fallback strategy (trusted-mint + # list), re-raise immediately so the caller can try the next + # mint instead of waiting through this mint's cooldown. + if not retry_on_rate_limit: + logger.warning( + "Mint rate-limited, skipping retries for fallback", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "cooldown_seconds": round(backoff, 2), + }, + ) + raise + if attempt >= max_attempts - 1: raise logger.warning( @@ -374,24 +406,13 @@ async def recieve_token( async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: - """Internal send function - returns amount and serialized token""" - effective_mint_url = mint_url or settings.primary_mint - wallet: Wallet = await get_wallet(effective_mint_url, unit) - proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) - proofs_for_mint = sum(p.amount for p in proofs) - - # Fallback: proofs from untrusted source mints are swapped to primary_mint - # during receive, so the user's preferred refund_mint_url may have no proofs - # even though the global wallet has the balance. - if proofs_for_mint < amount and effective_mint_url != settings.primary_mint: - logger.info( - f"send: insufficient proofs at {effective_mint_url} " - f"(have {proofs_for_mint}, need {amount}), falling back to primary_mint={settings.primary_mint}" - ) - effective_mint_url = settings.primary_mint - wallet = await get_wallet(effective_mint_url, unit) - proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) - proofs_for_mint = sum(p.amount for p in proofs) + """Create a token from the preferred mint or another funded trusted mint.""" + effective_mint_url = await find_trusted_mint_with_funds(amount, unit, mint_url) + wallet = await get_wallet(effective_mint_url, unit) + proofs = get_proofs_per_mint_and_unit( + wallet, effective_mint_url, unit, not_reserved=True + ) + proofs_for_mint = sum(proof.amount for proof in proofs) all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { @@ -435,6 +456,61 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str return token +async def release_token_reservation(token: str) -> None: + """Release a token that was created locally but never handed off.""" + token_obj = deserialize_token_from_string(token) + wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) + await wallet.set_reserved_for_send(token_obj.proofs, reserved=False) + + secrets = {proof.secret for proof in token_obj.proofs} + for proof in token_obj.proofs: + proof.reserved = False + for proof in wallet.proofs: + if proof.secret in secrets: + proof.reserved = False + + +def token_mint_url(token: str, fallback: str | None = None) -> str: + try: + return str(deserialize_token_from_string(token).mint) + except Exception: + if fallback is None: + raise + return fallback + + +async def find_trusted_mint_with_funds( + amount: int, unit: str, preferred_mint: str | None = None +) -> str: + """Choose a trusted mint that can cover a refund without waiting on cooldown.""" + trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) + candidates: list[str] = [] + if preferred_mint in trusted: + candidates.append(preferred_mint) + candidates.extend(mint for mint in trusted if mint not in candidates) + + balances: dict[str, int] = {} + for mint_url in candidates: + if _mint_cooldown_remaining(mint_url) > 0: + continue + try: + wallet = await get_wallet(mint_url, unit, retry_on_rate_limit=False) + except Exception as error: + if is_mint_connection_error(error) or _is_mint_rate_limited(error): + balances[mint_url] = 0 + continue + raise + + proofs = get_proofs_per_mint_and_unit(wallet, mint_url, unit, not_reserved=True) + balances[mint_url] = sum(proof.amount for proof in proofs) + if balances[mint_url] >= amount: + return mint_url + + raise ValueError( + f"No trusted mint has {amount} {unit} available; balances={balances}" + ) + + # A foreign mint's fee_reserve is a non-binding estimate (NUT-05): the mint may # demand more when re-quoting or at melt execution. Instead of padding the # estimate with a safety buffer (which strands the margin at the foreign mint @@ -505,21 +581,48 @@ async def _request_mint_with_fallback( amount: int, *, op_name: str, primary_wallet: Wallet | None = None ) -> tuple[Wallet, str, MintQuote]: """Try request_mint on the primary mint, fall back to other trusted mints - on transport or rate-limit failure. Returns the wallet, mint_url, and quote.""" + on transport or rate-limit failure. Returns the wallet, mint_url, and quote. + + Guards against amount <= 0: the cashu library's PostMintQuoteRequest + enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a + cryptic validation error deep in the stack. Fail fast with context. + """ + if amount <= 0: + raise ValueError( + f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " + f"Token value is too small after fee deduction or unit conversion." + ) candidates = [settings.primary_mint] + [ m for m in settings.cashu_mints if m != settings.primary_mint ] tried: list[str] = [] for mint_url in candidates: + cooldown = _mint_cooldown_remaining(mint_url) + if cooldown > 0: + tried.append(f"{mint_url}: cooling down") + logger.info( + "Skipping rate-limited mint", + extra={ + "mint_url": mint_url, + "cooldown_seconds": round(cooldown, 2), + "op_name": op_name, + }, + ) + continue try: if mint_url == settings.primary_mint and primary_wallet is not None: wallet = primary_wallet else: - wallet = await get_wallet(mint_url, settings.primary_mint_unit) + wallet = await get_wallet( + mint_url, + settings.primary_mint_unit, + retry_on_rate_limit=False, + ) quote = await _mint_operation( lambda: wallet.request_mint(amount), op_name=op_name, mint_url=mint_url, + retry_on_rate_limit=False, ) return wallet, mint_url, quote except Exception as e: @@ -563,11 +666,37 @@ async def _calculate_swap_amount( ) return int(receive_amount) + # The cashu library's PostMintQuoteRequest enforces amount > 0 (Pydantic + # Field(gt=0)). When the token's face value in the primary mint's unit + # truncates to 0 (e.g. < 1000 msat with a "sat" primary unit), calling + # request_mint(0) raises a validation error that is cryptic in production + # logs. Guard early with full diagnostic context instead. + if receive_amount <= 0: + logger.error( + "swap_to_primary_mint: receive_amount is zero or negative, cannot estimate fees", + extra={ + "amount_msat": amount_msat, + "token_unit": token_unit, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, + "primary_mint_unit": settings.primary_mint_unit, + "receive_amount": receive_amount, + }, + ) + raise ValueError( + f"Token amount ({amount_msat} msat, unit={token_unit}) is too small to " + f"swap to primary mint ({settings.primary_mint}, unit={settings.primary_mint_unit}): " + f"receive_amount={receive_amount}. Minimum 1 {settings.primary_mint_unit} required." + ) + logger.info( "swap_to_primary_mint: estimating fees", extra={ "dummy_amount": receive_amount, "unit": settings.primary_mint_unit, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, + "amount_msat": amount_msat, }, ) @@ -600,6 +729,9 @@ async def _calculate_swap_amount( "input_fees": input_fees, "minted_amount": minted_amount, "minted_unit": settings.primary_mint_unit, + "fee_reserve": fee_reserve, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, }, ) return minted_amount @@ -607,7 +739,16 @@ async def _calculate_swap_amount( except Exception as e: logger.error( "swap_to_primary_mint: fee estimation failed", - extra={"error": str(e)}, + extra={ + "error": str(e), + "error_type": type(e).__name__, + "amount_msat": amount_msat, + "token_unit": token_unit, + "token_mint_url": token_mint_url, + "primary_mint": settings.primary_mint, + "primary_mint_unit": settings.primary_mint_unit, + "receive_amount": receive_amount, + }, ) if is_mint_connection_error(e): raise MintConnectionError("Cashu mint is unreachable") from e @@ -672,6 +813,24 @@ async def swap_to_primary_mint( dest_mint_url = settings.primary_mint while True: attempt += 1 + if minted_amount <= 0: + logger.error( + "swap_to_primary_mint: minted_amount is zero or negative before requesting quote", + extra={ + "minted_amount": minted_amount, + "attempt": attempt, + "foreign_mint": token_obj.mint, + "token_amount": token_amount, + "token_unit": token_obj.unit, + "amount_msat": amount_msat, + "observed_extra_fee": observed_extra_fee, + "primary_mint": settings.primary_mint, + }, + ) + raise ValueError( + f"Cannot swap token ({token_amount} {token_obj.unit}) from {token_obj.mint}: " + f"minted_amount={minted_amount} after fee deduction (attempt {attempt})" + ) dest_wallet, dest_mint_url, mint_quote = await _request_mint_with_fallback( minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet ) @@ -1006,7 +1165,12 @@ _wallet_load_locks: dict[str, asyncio.Lock] = {} _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30 -async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet: +async def get_wallet( + mint_url: str, + unit: str = "sat", + load: bool = True, + retry_on_rate_limit: bool = True, +) -> Wallet: global _wallets, _wallet_last_load, _wallet_load_locks id = f"{mint_url}_{unit}" lock = _wallet_load_locks.setdefault(id, asyncio.Lock()) @@ -1016,17 +1180,19 @@ async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wal if load: now = time.monotonic() - last = _wallet_last_load.get(id, 0) - if now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: + last = _wallet_last_load.get(id) + if last is None or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: await _mint_operation( lambda: _wallets[id].load_mint(), op_name="load_mint", mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, ) await _mint_operation( lambda: _wallets[id].load_proofs(reload=True), op_name="load_proofs", mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, ) _wallet_last_load[id] = time.monotonic() return _wallets[id] @@ -1361,9 +1527,10 @@ async def periodic_routstr_fee_payout() -> None: async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: + mint = await find_trusted_mint_with_funds(amount, unit, mint) wallet = await get_wallet(mint, unit) - proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id] - proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) + available = get_proofs_per_mint_and_unit(wallet, mint, unit, not_reserved=True) + proofs, _ = await wallet.select_to_send(available, amount, set_reserved=True) return await raw_send_to_lnurl(wallet, proofs, address, unit) diff --git a/tests/unit/test_auto_topup.py b/tests/unit/test_auto_topup.py index c05c85e5..cf3d06e2 100644 --- a/tests/unit/test_auto_topup.py +++ b/tests/unit/test_auto_topup.py @@ -66,6 +66,10 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected() "routstr.upstream.auto_topup.store_cashu_transaction", AsyncMock(return_value=True), ) as store, + patch( + "routstr.upstream.auto_topup.token_mint_url", + return_value="https://fallback-mint.test", + ), patch("routstr.upstream.auto_topup.create_session", return_value=session), ): await _check_and_topup(_row()) @@ -74,7 +78,7 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected() token="cashu-token", amount=50, unit="sat", - mint_url="https://mint.test", + mint_url="https://fallback-mint.test", typ="out", collected=False, source="auto_topup", @@ -138,6 +142,12 @@ async def test_auto_topup_does_not_send_untracked_token() -> None: "routstr.upstream.auto_topup.store_cashu_transaction", AsyncMock(return_value=False), ), + patch( + "routstr.upstream.auto_topup.release_token_reservation", + AsyncMock(), + ) as reclaim, ): await _check_and_topup(_row()) + + reclaim.assert_awaited_once_with("cashu-token") provider.topup.assert_not_awaited() diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index f712b357..376442d5 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -169,6 +169,59 @@ async def test_send_token() -> None: assert token == "test_token" +@pytest.mark.asyncio +async def test_release_token_reservation_unreserves_local_proofs() -> None: + from routstr.wallet import release_token_reservation + + token_proof = Mock(secret="proof-secret", reserved=True) + cached_proof = Mock(secret="proof-secret", reserved=True) + token = Mock(mint="http://mint:3338", unit="sat", proofs=[token_proof]) + wallet = Mock(proofs=[cached_proof], set_reserved_for_send=AsyncMock()) + with ( + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch( + "routstr.wallet.get_wallet", AsyncMock(return_value=wallet) + ) as get_wallet, + ): + await release_token_reservation("cashu-token") + + get_wallet.assert_awaited_once_with("http://mint:3338", "sat", load=False) + wallet.set_reserved_for_send.assert_awaited_once_with(token.proofs, reserved=False) + assert token_proof.reserved is False + assert cached_proof.reserved is False + + +@pytest.mark.asyncio +async def test_refund_mint_falls_back_to_trusted_mint_with_funds() -> None: + from routstr.core.settings import settings + from routstr.wallet import find_trusted_mint_with_funds + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + def wallet_for(mint: str, amount: int) -> Mock: + keyset = Mock(id=f"keyset-{mint}", mint_url=mint) + keyset.unit.name = "sat" + proof = Mock(id=keyset.id, amount=amount, reserved=False) + return Mock(keysets={keyset.id: keyset}, proofs=[proof]) + + wallets = { + primary: wallet_for(primary, 50), + secondary: wallet_for(secondary, 200), + } + with ( + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary, secondary]), + patch( + "routstr.wallet.get_wallet", + AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), + ), + ): + mint = await find_trusted_mint_with_funds(100, "sat", primary) + + assert mint == secondary + + @pytest.mark.asyncio async def test_credit_balance() -> None: token_data = { @@ -1416,6 +1469,25 @@ def test_mint_rate_guard_rebuilds_when_setting_changes() -> None: assert second._max_concurrency == 2 +@pytest.mark.asyncio +async def test_mint_rate_guard_keeps_cooldown_when_concurrency_is_unlimited() -> None: + from routstr.core.settings import settings + from routstr.wallet import _MintRateGuard + + operation = AsyncMock(return_value="ok") + with ( + patch.object(settings, "mint_max_concurrency", 0), + patch("routstr.wallet.time.monotonic", return_value=0), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + ): + guard = _MintRateGuard.get("http://mint:3338") + guard.apply_cooldown(5) + assert await guard.run(operation) == "ok" + + sleep.assert_awaited_once_with(5) + operation.assert_awaited_once() + + @pytest.mark.asyncio async def test_mint_operation_honors_retry_after_as_minimum() -> None: from routstr.core.settings import settings @@ -1495,7 +1567,9 @@ async def test_get_wallet_initializes_and_loads_once_concurrently() -> None: with patch( "routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet) ) as create: - with patch("routstr.wallet.time.monotonic", return_value=100.0): + # A fresh wallet must load even when the host has been up for less than + # the reload interval. + with patch("routstr.wallet.time.monotonic", return_value=10.0): first, second = await asyncio.gather( get_wallet("http://mint:3338"), get_wallet("http://mint:3338") ) @@ -1506,6 +1580,36 @@ async def test_get_wallet_initializes_and_loads_once_concurrently() -> None: mock_wallet.load_proofs.assert_awaited_once_with(reload=True) +@pytest.mark.asyncio +async def test_get_wallet_can_surface_429_without_retrying() -> None: + from routstr.core.settings import settings + from routstr.wallet import get_wallet + + request = httpx.Request("GET", "http://mint:3338/v1/info") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + wallet = Mock( + load_mint=AsyncMock( + side_effect=httpx.HTTPStatusError( + "rate limited", request=request, response=response + ) + ), + load_proofs=AsyncMock(), + ) + + with ( + patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=wallet)), + patch.object(settings, "mint_retry_max_attempts", 3), + patch.object(settings, "mint_operation_timeout_seconds", 0), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + ): + with pytest.raises(httpx.HTTPStatusError): + await get_wallet("http://mint:3338", retry_on_rate_limit=False) + + wallet.load_mint.assert_awaited_once() + wallet.load_proofs.assert_not_awaited() + sleep.assert_not_awaited() + + @pytest.mark.asyncio async def test_mint_operation_factory_retry_succeeds() -> None: """_mint_operation accepts a zero-arg factory, not a dead coroutine. @@ -1752,3 +1856,161 @@ async def test_lightning_mint_fallback_all_fail() -> None: ): with pytest.raises(MintConnectionError): await _request_mint_with_fallback(1000) + + +@pytest.mark.asyncio +async def test_lightning_mint_fallback_rejects_zero_amount() -> None: + """Zero or negative amounts must be rejected before reaching the mint.""" + from routstr.lightning import _request_mint_with_fallback + + with pytest.raises(ValueError, match="amount_sats must be > 0"): + await _request_mint_with_fallback(0) + + with pytest.raises(ValueError, match="amount_sats must be > 0"): + await _request_mint_with_fallback(-5) + + +@pytest.mark.asyncio +async def test_wallet_request_mint_fallback_rejects_zero_amount() -> None: + """Zero or negative amounts must be rejected before reaching the mint.""" + from routstr.wallet import _request_mint_with_fallback + + with pytest.raises(ValueError, match="amount must be > 0"): + await _request_mint_with_fallback(0, op_name="test") + + with pytest.raises(ValueError, match="amount must be > 0"): + await _request_mint_with_fallback(-1, op_name="test") + + +@pytest.mark.asyncio +async def test_wallet_fallback_on_429_no_in_place_retry() -> None: + """A 429 from the primary mint must trigger immediate fallback to the + secondary — _mint_operation must NOT retry in-place when + retry_on_rate_limit=False is set by _request_mint_with_fallback.""" + from routstr.core.settings import settings + from routstr.wallet import _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + request = httpx.Request("POST", "http://primary:3338/v1/mint/quote/bolt11") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + primary_call_count = 0 + + async def primary_request_mint(_amount: int) -> None: + nonlocal primary_call_count + primary_call_count += 1 + raise httpx.HTTPStatusError("rate limited", request=request, response=response) + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock(side_effect=primary_request_mint) + + mock_quote = Mock(quote="q_secondary", request="lnbc1secondary") + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_retry_max_attempts", 3): + with patch.object(settings, "mint_max_concurrency", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("asyncio.sleep", AsyncMock()) as mock_sleep: + with patch( + "routstr.wallet.get_wallet", side_effect=mock_get + ): + _, mint_url, _ = await _request_mint_with_fallback( + 1000, op_name="test_429_fallback" + ) + + assert mint_url == secondary + assert primary_call_count == 1 + mock_secondary_wallet.request_mint.assert_called_once() + mock_sleep.assert_not_called() + + +@pytest.mark.asyncio +async def test_wallet_fallback_skips_mint_during_cooldown() -> None: + from routstr.core.settings import settings + from routstr.wallet import _MintRateGuard, _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + primary_wallet = Mock(request_mint=AsyncMock()) + quote = Mock(quote="q_secondary", request="lnbc1secondary") + secondary_wallet = Mock(request_mint=AsyncMock(return_value=quote)) + wallets = {primary: primary_wallet, secondary: secondary_wallet} + + with ( + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary, secondary]), + patch.object(settings, "mint_max_concurrency", 0), + patch.object(settings, "mint_operation_timeout_seconds", 0), + patch("routstr.wallet.time.monotonic", return_value=10), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + patch( + "routstr.wallet.get_wallet", + AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), + ), + ): + _MintRateGuard.get(primary).apply_cooldown(60) + _, mint_url, _ = await _request_mint_with_fallback( + 1000, op_name="test_cooldown_fallback" + ) + + assert mint_url == secondary + primary_wallet.request_mint.assert_not_awaited() + secondary_wallet.request_mint.assert_awaited_once_with(1000) + sleep.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_lightning_fallback_on_429_no_in_place_retry() -> None: + """Same as above but for the lightning.py _request_mint_with_fallback.""" + from routstr.core.settings import settings + from routstr.lightning import _request_mint_with_fallback + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + + request = httpx.Request("POST", "http://primary:3338/v1/mint/quote/bolt11") + response = httpx.Response(429, request=request, headers={"Retry-After": "60"}) + primary_call_count = 0 + + async def primary_request_mint(_amount: int) -> None: + nonlocal primary_call_count + primary_call_count += 1 + raise httpx.HTTPStatusError("rate limited", request=request, response=response) + + mock_primary_wallet = Mock() + mock_primary_wallet.request_mint = AsyncMock(side_effect=primary_request_mint) + + mock_quote = Mock(quote="q_secondary", request="lnbc1secondary") + mock_secondary_wallet = Mock() + mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote) + + wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} + mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + + with patch.object(settings, "primary_mint", primary): + with patch.object(settings, "cashu_mints", [primary, secondary]): + with patch.object(settings, "mint_retry_max_attempts", 3): + with patch.object(settings, "mint_max_concurrency", 0): + with patch.object(settings, "mint_operation_timeout_seconds", 0): + with patch("asyncio.sleep", AsyncMock()) as mock_sleep: + with patch( + "routstr.lightning.get_wallet", side_effect=mock_get + ): + _, _, first_mint = await _request_mint_with_fallback( + 1000 + ) + _, _, second_mint = await _request_mint_with_fallback( + 1000 + ) + + assert first_mint == second_mint == secondary + assert primary_call_count == 1 + assert mock_secondary_wallet.request_mint.await_count == 2 + mock_sleep.assert_not_called() From d44b98fd0dd268aebe3b5727ae1ee5cd9cbd3daa Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 14 Jul 2026 00:54:44 +0200 Subject: [PATCH 12/46] fix fallback --- routstr/balance.py | 58 ++++- routstr/wallet.py | 283 +++++++++++++++++++---- tests/integration/test_swap_fee_retry.py | 7 +- tests/unit/test_fetch_all_balances.py | 49 +++- tests/unit/test_wallet.py | 105 ++++----- 5 files changed, 394 insertions(+), 108 deletions(-) diff --git a/routstr/balance.py b/routstr/balance.py index 8332577d..7739268f 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -145,6 +145,17 @@ class TopupRequest(BaseModel): cashu_token: str +def _error_chain(error: BaseException) -> list[dict[str, str]]: + chain: list[dict[str, str]] = [] + current: BaseException | None = error + seen: set[int] = set() + while current is not None and id(current) not in seen: + seen.add(id(current)) + chain.append({"type": type(current).__name__, "message": str(current)}) + current = current.__cause__ or current.__context__ + return chain + + @router.post("/topup") async def topup_wallet_endpoint( cashu_token: str | None = None, @@ -162,6 +173,18 @@ async def topup_wallet_endpoint( cashu_token = cashu_token.replace("\n", "").replace("\r", "").replace("\t", "") if len(cashu_token) < 10 or "cashu" not in cashu_token: raise HTTPException(status_code=400, detail="Invalid token format") + + source_mint = token_mint_url(cashu_token, "unknown") + logger.warning( + "Cashu wallet top-up started", + extra={ + "event": "cashu_topup_started", + "source_mint": source_mint, + "primary_mint": settings.primary_mint, + "trusted_mints": settings.cashu_mints, + "key_hash": billing_key.hashed_key[:8], + }, + ) try: amount_msats = await credit_balance(cashu_token, billing_key, session) except Exception as e: @@ -170,12 +193,41 @@ async def topup_wallet_endpoint( classified = classify_redemption_error(e) if classified is None: logger.error( - "topup_wallet_endpoint: unhandled error", - extra={"error": str(e), "error_type": type(e).__name__}, + "Cashu wallet top-up failed with an unhandled error", + extra={ + "event": "cashu_topup_failed", + "source_mint": source_mint, + "primary_mint": settings.primary_mint, + "trusted_mints": settings.cashu_mints, + "error_chain": _error_chain(e), + }, ) raise HTTPException(status_code=500, detail="Internal server error") - _type, status_code, message, _code = classified + error_type, status_code, message, error_code = classified + logger.warning( + "Cashu wallet top-up failed", + extra={ + "event": "cashu_topup_failed", + "source_mint": source_mint, + "primary_mint": settings.primary_mint, + "trusted_mints": settings.cashu_mints, + "status_code": status_code, + "error_type": error_type, + "error_code": error_code, + "error_chain": _error_chain(e), + }, + ) raise HTTPException(status_code=status_code, detail=message) + + logger.warning( + "Cashu wallet top-up completed", + extra={ + "event": "cashu_topup_completed", + "source_mint": source_mint, + "credited_msats": amount_msats, + "key_hash": billing_key.hashed_key[:8], + }, + ) return {"msats": amount_msats} diff --git a/routstr/wallet.py b/routstr/wallet.py index f992c7a0..b6ae1bcb 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -53,6 +53,8 @@ class TokenConsumedError(Exception): # httpx base classes cover their subclasses. HTTPStatusError is excluded on # purpose — that means the mint answered, just with an error status. +_MINT_TRANSPORT_COOLDOWN_SECONDS = 30.0 + _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( httpx.NetworkError, httpx.TimeoutException, @@ -400,8 +402,30 @@ async def recieve_token( wallet.keyset_id = token_obj.keysets[0] if token_obj.mint not in settings.cashu_mints: + destinations = list( + dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) + ) + logger.warning( + "Cashu cross-mint swap required", + extra={ + "event": "cashu_swap_started", + "source_mint": token_obj.mint, + "source_unit": token_obj.unit, + "source_amount": token_obj.amount, + "destination_candidates": destinations, + }, + ) return await swap_to_primary_mint(token_obj, wallet) + logger.info( + "Cashu same-mint redemption selected", + extra={ + "event": "cashu_same_mint_redemption", + "source_mint": token_obj.mint, + "source_unit": token_obj.unit, + "source_amount": token_obj.amount, + }, + ) return await _redeem_same_mint(wallet, token_obj) @@ -595,20 +619,43 @@ async def _request_mint_with_fallback( candidates = [settings.primary_mint] + [ m for m in settings.cashu_mints if m != settings.primary_mint ] + logger.warning( + "Trying trusted destination mints", + extra={ + "event": "cashu_destination_candidates", + "op_name": op_name, + "amount": amount, + "unit": settings.primary_mint_unit, + "candidates": candidates, + }, + ) tried: list[str] = [] - for mint_url in candidates: + for candidate_index, mint_url in enumerate(candidates, start=1): cooldown = _mint_cooldown_remaining(mint_url) if cooldown > 0: tried.append(f"{mint_url}: cooling down") - logger.info( - "Skipping rate-limited mint", + logger.warning( + "Skipping unavailable destination mint", extra={ + "event": "cashu_destination_skipped", "mint_url": mint_url, "cooldown_seconds": round(cooldown, 2), "op_name": op_name, + "candidate_index": candidate_index, + "candidate_count": len(candidates), }, ) continue + logger.warning( + "Trying destination mint", + extra={ + "event": "cashu_destination_attempt", + "mint_url": mint_url, + "op_name": op_name, + "candidate_index": candidate_index, + "candidate_count": len(candidates), + }, + ) try: if mint_url == settings.primary_mint and primary_wallet is not None: wallet = primary_wallet @@ -624,21 +671,54 @@ async def _request_mint_with_fallback( mint_url=mint_url, retry_on_rate_limit=False, ) - return wallet, mint_url, quote - except Exception as e: - tried.append(f"{mint_url}: {type(e).__name__}") - if not is_mint_connection_error(e) and not _is_mint_rate_limited(e): - raise logger.warning( - "request_mint failed, trying fallback mint", + "Destination mint selected", extra={ + "event": "cashu_destination_selected", + "mint_url": mint_url, + "op_name": op_name, + "candidate_index": candidate_index, + "fallback_used": candidate_index > 1, + }, + ) + return wallet, mint_url, quote + except Exception as error: + tried.append(f"{mint_url}: {type(error).__name__}") + connection_failure = is_mint_connection_error(error) + rate_limited = _is_mint_rate_limited(error) + if not connection_failure and not rate_limited: + raise + if connection_failure: + _MintRateGuard.get(mint_url).apply_cooldown( + _MINT_TRANSPORT_COOLDOWN_SECONDS + ) + logger.warning( + "Destination mint failed", + extra={ + "event": "cashu_destination_failed", "failed_mint": mint_url, - "error": str(e), + "error": str(error), + "error_type": type(error).__name__, + "connection_failure": connection_failure, + "rate_limited": rate_limited, "tried": tried, "op_name": op_name, + "candidate_index": candidate_index, + "candidate_count": len(candidates), }, ) continue + logger.error( + "All trusted destination mints failed", + extra={ + "event": "cashu_destination_exhausted", + "op_name": op_name, + "amount": amount, + "unit": settings.primary_mint_unit, + "candidates": candidates, + "tried": tried, + }, + ) raise MintConnectionError(f"All mints failed for {op_name}: {tried}") @@ -647,7 +727,7 @@ async def _calculate_swap_amount( token_unit: str, token_mint_url: str, token_wallet: Wallet, - primary_wallet: Wallet, + primary_wallet: Wallet | None, proofs: list, ) -> int: """ @@ -700,12 +780,14 @@ async def _calculate_swap_amount( }, ) + stage = "destination_fee_quote" try: _, _, dummy_mint_quote = await _request_mint_with_fallback( receive_amount, op_name="swap_fee_est_mint_quote", primary_wallet=primary_wallet, ) + stage = "source_fee_quote" dummy_melt_quote = await _mint_operation( lambda: token_wallet.melt_quote(dummy_mint_quote.request), op_name="swap_fee_est_melt_quote", @@ -738,8 +820,10 @@ async def _calculate_swap_amount( except Exception as e: logger.error( - "swap_to_primary_mint: fee estimation failed", + "Cashu swap fee estimation failed", extra={ + "event": "cashu_swap_fee_estimation_failed", + "stage": stage, "error": str(e), "error_type": type(e).__name__, "amount_msat": amount_msat, @@ -751,6 +835,15 @@ async def _calculate_swap_amount( }, ) if is_mint_connection_error(e): + if stage == "source_fee_quote": + logger.error( + "Source mint is unreachable; destination fallback cannot spend its proofs", + extra={ + "event": "cashu_source_mint_unreachable", + "source_mint": token_mint_url, + "stage": stage, + }, + ) raise MintConnectionError("Cashu mint is unreachable") from e raise ValueError(f"Failed to estimate fees: {e}") from e @@ -758,10 +851,11 @@ async def _calculate_swap_amount( async def swap_to_primary_mint( token_obj: Token, token_wallet: Wallet ) -> tuple[int, str, str]: - logger.info( - "swap_to_primary_mint: starting", + logger.warning( + "Starting Cashu cross-mint swap", extra={ - "foreign_mint": token_obj.mint, + "event": "cashu_swap_started", + "source_mint": token_obj.mint, "token_amount": token_obj.amount, "unit": token_obj.unit, "primary_mint": settings.primary_mint, @@ -793,7 +887,7 @@ async def swap_to_primary_mint( ) return await _redeem_same_mint(token_wallet, token_obj) - primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit) + primary_wallet: Wallet | None = None minted_amount = await _calculate_swap_amount( amount_msat, @@ -843,11 +937,37 @@ async def swap_to_primary_mint( }, ) - melt_quote = await _mint_operation( - lambda: token_wallet.melt_quote(mint_quote.request), - op_name="swap_melt_quote", - mint_url=token_obj.mint, + logger.warning( + "Requesting melt quote from source mint", + extra={ + "event": "cashu_source_melt_quote_attempt", + "source_mint": token_obj.mint, + "destination_mint": dest_mint_url, + "attempt": attempt, + }, ) + try: + melt_quote = await _mint_operation( + lambda: token_wallet.melt_quote(mint_quote.request), + op_name="swap_melt_quote", + mint_url=token_obj.mint, + ) + except Exception as error: + if is_mint_connection_error(error): + logger.error( + "Source mint is unreachable; destination fallback cannot spend its proofs", + extra={ + "event": "cashu_source_mint_unreachable", + "source_mint": token_obj.mint, + "destination_mint": dest_mint_url, + "stage": "source_melt_quote", + "error": str(error), + "error_type": type(error).__name__, + "attempt": attempt, + }, + ) + raise MintConnectionError("Cashu mint is unreachable") from error + raise input_fees = token_wallet.get_fees_for_proofs(token_obj.proofs) total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees logger.info( @@ -915,8 +1035,16 @@ async def swap_to_primary_mint( # A down mint won't fix itself by retrying with a smaller amount. if is_mint_connection_error(e): logger.error( - "swap_to_primary_mint: melt failed — mint unreachable", - extra={"error": str(e), "foreign_mint": token_obj.mint}, + "Source mint became unreachable during melt", + extra={ + "event": "cashu_source_mint_unreachable", + "stage": "source_melt", + "error": str(e), + "error_type": type(e).__name__, + "source_mint": token_obj.mint, + "destination_mint": dest_mint_url, + "attempt": attempt, + }, ) raise MintConnectionError("Cashu mint is unreachable") from e shortfall = _melt_insufficient_shortfall(e) @@ -957,9 +1085,10 @@ async def swap_to_primary_mint( break - logger.info( - "swap_to_primary_mint: melt succeeded, minting on destination", + logger.warning( + "Source melt succeeded; minting on destination", extra={ + "event": "cashu_destination_mint_attempt", "minted_amount": minted_amount, "mint_quote_id": mint_quote.quote, "dest_mint": dest_mint_url, @@ -1040,10 +1169,11 @@ async def swap_to_primary_mint( "Mint on primary failed after successful melt" ) from e - logger.info( - "swap_to_primary_mint: completed successfully", + logger.warning( + "Cashu cross-mint swap completed", extra={ - "foreign_mint": token_obj.mint, + "event": "cashu_swap_completed", + "source_mint": token_obj.mint, "dest_mint": dest_mint_url, "original_amount": token_amount, "minted_amount": minted_amount, @@ -1058,8 +1188,11 @@ async def credit_balance( cashu_token: str, key: db.ApiKey, session: db.AsyncSession ) -> int: logger.info( - "credit_balance: Starting token redemption", - extra={"token_preview": cashu_token[:50]}, + "Starting Cashu balance credit", + extra={ + "event": "cashu_credit_started", + "key_hash": key.hashed_key[:8], + }, ) try: @@ -1212,7 +1345,12 @@ def get_proofs_per_mint_and_unit( return proofs -async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[Proof]: +async def slow_filter_spend_proofs( + proofs: list[Proof], + wallet: Wallet, + *, + retry_on_rate_limit: bool = True, +) -> list[Proof]: if not proofs: return [] _proofs = [] @@ -1226,6 +1364,7 @@ async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[ lambda: wallet.check_proof_state(pb), op_name="check_proof_state", mint_url=str(wallet.url), + retry_on_rate_limit=retry_on_rate_limit, ) for proof, state in zip(pb, proof_states.states): if str(state.state) != "spent": @@ -1251,6 +1390,22 @@ class BalanceDetail(TypedDict, total=False): error: str +_BALANCE_FETCH_RETRY_SECONDS = 60.0 +_balance_fetch_failures: dict[tuple[str, str], tuple[float, str]] = {} +_balance_fetch_locks: dict[tuple[str, str], asyncio.Lock] = {} + + +def _balance_error(mint_url: str, unit: str, error: str) -> BalanceDetail: + return { + "mint_url": mint_url, + "unit": unit, + "wallet_balance": 0, + "user_balance": 0, + "owner_balance": 0, + "error": error, + } + + async def fetch_all_balances( units: list[str] | None = None, ) -> tuple[list[BalanceDetail], int, int, int]: @@ -1269,18 +1424,56 @@ async def fetch_all_balances( async def fetch_balance( session: db.AsyncSession, mint_url: str, unit: str ) -> BalanceDetail: - try: - wallet = await get_wallet(mint_url, unit) - proofs = get_proofs_per_mint_and_unit( - wallet, mint_url, unit, not_reserved=True - ) - proofs = await slow_filter_spend_proofs(proofs, wallet) - user_balance = await db.balances_for_mint_and_unit(session, mint_url, unit) + key = (mint_url, unit) + lock = _balance_fetch_locks.setdefault(key, asyncio.Lock()) + async with lock: + now = time.monotonic() + failure = _balance_fetch_failures.get(key) + if failure is not None and now < failure[0]: + return _balance_error(mint_url, unit, failure[1]) + + cooldown = _mint_cooldown_remaining(mint_url) + if cooldown > 0: + error = "Mint is cooling down after a rate limit" + _balance_fetch_failures[key] = (now + cooldown, error) + return _balance_error(mint_url, unit, error) + + try: + wallet = await get_wallet( + mint_url, unit, retry_on_rate_limit=False + ) + proofs = get_proofs_per_mint_and_unit( + wallet, mint_url, unit, not_reserved=True + ) + proofs = await slow_filter_spend_proofs( + proofs, wallet, retry_on_rate_limit=False + ) + user_balance = await db.balances_for_mint_and_unit( + session, mint_url, unit + ) + except Exception as error: + retry_delay = max( + _BALANCE_FETCH_RETRY_SECONDS, + _mint_cooldown_remaining(mint_url), + ) + retry_at = time.monotonic() + retry_delay + _balance_fetch_failures[key] = (retry_at, str(error)) + logger.warning( + "Unable to refresh mint balance", + extra={ + "mint_url": mint_url, + "unit": unit, + "error": str(error), + "retry_seconds": round(retry_delay, 2), + }, + ) + return _balance_error(mint_url, unit, str(error)) + + _balance_fetch_failures.pop(key, None) if unit == "sat": user_balance = user_balance // 1000 proofs_balance = sum(proof.amount for proof in proofs) - - result: BalanceDetail = { + return { "mint_url": mint_url, "unit": unit, "wallet_balance": proofs_balance, @@ -1289,18 +1482,6 @@ async def fetch_all_balances( if proofs_balance != 0 else 0, } - return result - except Exception as e: - logger.error(f"Error getting balance for {mint_url} {unit}: {e}") - error_result: BalanceDetail = { - "mint_url": mint_url, - "unit": unit, - "wallet_balance": 0, - "user_balance": 0, - "owner_balance": 0, - "error": str(e), - } - return error_result # Build the set of mints to inspect. Received tokens are stored against # ``primary_mint`` (which defaults to a real mint even when ``cashu_mints`` diff --git a/tests/integration/test_swap_fee_retry.py b/tests/integration/test_swap_fee_retry.py index d2326a89..138a4d4d 100644 --- a/tests/integration/test_swap_fee_retry.py +++ b/tests/integration/test_swap_fee_retry.py @@ -89,7 +89,12 @@ def _make_swap_mocks( def _wallet_router(primary_wallet: Mock, token_wallet: Mock) -> Callable[..., Mock]: """Route get_wallet calls to the primary or foreign wallet mock by URL.""" - def fake_get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Mock: + def fake_get_wallet( + mint_url: str, + unit: str = "sat", + load: bool = True, + **kwargs: object, + ) -> Mock: return primary_wallet if mint_url == PRIMARY_MINT else token_wallet return fake_get_wallet diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index dcd99107..20cbe312 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -1,11 +1,24 @@ +from collections.abc import Generator from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from routstr.wallet import fetch_all_balances +@pytest.fixture(autouse=True) +def clear_balance_fetch_state() -> Generator[None, None, None]: + from routstr import wallet + + wallet._balance_fetch_failures.clear() + wallet._balance_fetch_locks.clear() + yield + wallet._balance_fetch_failures.clear() + wallet._balance_fetch_locks.clear() + + @asynccontextmanager async def _fake_session(): # type: ignore[no-untyped-def] yield MagicMock() @@ -21,7 +34,7 @@ def _patches(proof_amount: int = 1000): # type: ignore[no-untyped-def] ), patch( "routstr.wallet.slow_filter_spend_proofs", - AsyncMock(side_effect=lambda proofs, wallet: proofs), + AsyncMock(side_effect=lambda proofs, wallet, **kwargs: proofs), ), patch( "routstr.wallet.db.balances_for_mint_and_unit", @@ -52,6 +65,40 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None: assert total_wallet == 1000 +@pytest.mark.asyncio +async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: + from routstr.core.settings import settings + + get_wallet = AsyncMock(side_effect=httpx.ConnectError("mint unavailable")) + with ( + patch.object(settings, "cashu_mints", ["http://mint:3338"]), + patch.object(settings, "primary_mint", "http://mint:3338"), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.db.create_session", _fake_session), + patch("routstr.wallet.time.monotonic", return_value=10), + patch("routstr.wallet.logger.warning") as warning, + ): + first = await fetch_all_balances(units=["sat"]) + second = await fetch_all_balances(units=["sat"]) + + assert first[0][0]["error"] == "mint unavailable" + assert second[0][0]["error"] == "mint unavailable" + assert get_wallet.await_count == 1 + warning.assert_called_once() + + with ( + patch.object(settings, "cashu_mints", ["http://mint:3338"]), + patch.object(settings, "primary_mint", "http://mint:3338"), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.db.create_session", _fake_session), + patch("routstr.wallet.time.monotonic", return_value=71), + patch("routstr.wallet.logger.warning"), + ): + await fetch_all_balances(units=["sat"]) + + assert get_wallet.await_count == 2 + + @pytest.mark.asyncio async def test_fetch_all_balances_no_duplicate_primary_mint() -> None: """primary_mint already in cashu_mints is not inspected twice.""" diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 376442d5..ff666e65 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1710,71 +1710,72 @@ async def test_lightning_mint_fallback_for_topups() -> None: @pytest.mark.asyncio -async def test_swap_falls_back_to_secondary_mint() -> None: - """When the primary mint is unreachable, swap_to_primary_mint falls back - to a secondary trusted mint as the swap destination.""" +async def test_swap_falls_back_when_primary_wallet_cannot_load() -> None: from routstr.core.settings import settings - from routstr.wallet import _wallet_last_load, _wallets, swap_to_primary_mint - - _wallets.clear() - _wallet_last_load.clear() + from routstr.wallet import swap_to_primary_mint primary = "http://primary:3338" secondary = "http://secondary:3338" foreign = "http://foreign:3338" - mock_token = Mock() - mock_token.mint = foreign - mock_token.unit = "sat" - mock_token.amount = 1000 - mock_token.keysets = ["keyset1"] - mock_token.proofs = [Mock(amount=1000)] - - mock_token_wallet = Mock() - mock_token_wallet.load_mint = AsyncMock() - mock_token_wallet.load_proofs = AsyncMock() - mock_token_wallet.get_fees_for_proofs = Mock(return_value=0) - mock_token_wallet.melt_quote = AsyncMock( - return_value=Mock(quote="melt_q", amount=990, fee_reserve=10) + token = Mock( + mint=foreign, + unit="sat", + amount=1000, + keysets=["keyset1"], + proofs=[Mock(amount=1000)], ) - mock_token_wallet.melt = AsyncMock(return_value=Mock()) - - mock_primary_wallet = Mock() - mock_primary_wallet.request_mint = AsyncMock( - side_effect=httpx.ConnectError("primary down") + source_wallet = Mock( + load_mint=AsyncMock(), + load_proofs=AsyncMock(), + get_fees_for_proofs=Mock(return_value=0), + melt_quote=AsyncMock( + return_value=Mock(quote="melt_q", amount=990, fee_reserve=10) + ), + melt=AsyncMock(return_value=Mock()), ) mint_quote = Mock(quote="mint_q_secondary", request="lnbc1secondary") - mock_secondary_wallet = Mock() - mock_secondary_wallet.load_mint = AsyncMock() - mock_secondary_wallet.load_proofs = AsyncMock() - mock_secondary_wallet.available_balance = Mock(amount=0) - mock_secondary_wallet.keysets = ["ks_secondary"] - mock_secondary_wallet.restore_tokens_for_keyset = AsyncMock() - mock_secondary_wallet.request_mint = AsyncMock(return_value=mint_quote) - mock_secondary_wallet.mint = AsyncMock(return_value=Mock()) + secondary_wallet = Mock( + load_mint=AsyncMock(), + load_proofs=AsyncMock(), + available_balance=Mock(amount=0), + keysets=["ks_secondary"], + restore_tokens_for_keyset=AsyncMock(), + request_mint=AsyncMock(return_value=mint_quote), + mint=AsyncMock(return_value=Mock()), + ) - wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet} - mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m]) + async def get_wallet(mint: str, *args: object, **kwargs: object) -> Mock: + if mint == primary: + raise httpx.ConnectError("primary down") + return secondary_wallet - with patch.object(settings, "primary_mint", primary): - with patch.object(settings, "primary_mint_unit", "sat"): - with patch.object(settings, "cashu_mints", [primary, secondary]): - with patch.object(settings, "mint_max_concurrency", 0): - with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch("asyncio.sleep", AsyncMock()): - with patch( - "routstr.wallet.get_wallet", side_effect=mock_get - ): - amount, unit, mint_url = await swap_to_primary_mint( - mock_token, mock_token_wallet - ) + mock_get = AsyncMock(side_effect=get_wallet) + with ( + patch.object(settings, "primary_mint", primary), + patch.object(settings, "primary_mint_unit", "sat"), + patch.object(settings, "cashu_mints", [primary, secondary]), + patch.object(settings, "mint_max_concurrency", 0), + patch.object(settings, "mint_operation_timeout_seconds", 0), + patch("asyncio.sleep", AsyncMock()), + patch("routstr.wallet.get_wallet", side_effect=mock_get), + patch("routstr.wallet.logger.warning") as warning, + ): + amount, unit, mint_url = await swap_to_primary_mint(token, source_wallet) - assert mint_url == secondary - assert amount == 990 # 1000 - 10 fee_reserve - assert unit == "sat" - mock_secondary_wallet.mint.assert_called_once() - mock_primary_wallet.mint.assert_not_called() + assert (amount, unit, mint_url) == (990, "sat", secondary) + secondary_wallet.mint.assert_awaited_once() + assert mock_get.await_args_list[0].args[0] == primary + assert any(call.args[0] == secondary for call in mock_get.await_args_list) + events = { + call.kwargs["extra"]["event"] + for call in warning.call_args_list + if "extra" in call.kwargs and "event" in call.kwargs["extra"] + } + assert "cashu_destination_failed" in events + assert "cashu_destination_selected" in events + assert "cashu_swap_completed" in events @pytest.mark.asyncio From d7c401d2048262652516827878cb8dee52c37935 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 14 Jul 2026 01:10:23 +0200 Subject: [PATCH 13/46] primary mint fallback --- routstr/wallet.py | 181 +++++++++++++++++++++++++++++++------ tests/unit/test_balance.py | 23 +++++ tests/unit/test_wallet.py | 60 +++++++++++- 3 files changed, 236 insertions(+), 28 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index b6ae1bcb..a898ab7c 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -40,6 +40,10 @@ class MintConnectionError(Exception): """ +class SourceMintConnectionError(MintConnectionError): + """The mint that issued the incoming proofs cannot be reached.""" + + class TokenConsumedError(Exception): """A failure that happened AFTER the token's proofs were spent (melt succeeded, or redemption already returned) — e.g. minting on the primary @@ -245,6 +249,17 @@ def _parse_retry_after(headers: Any) -> float | None: return None +def is_source_mint_connection_error(error: BaseException) -> bool: + seen: set[int] = set() + current: BaseException | None = error + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, SourceMintConnectionError): + return True + current = current.__cause__ or current.__context__ + return False + + def is_mint_connection_error(error: BaseException) -> bool: """True if ``error`` (or anything in its cause/context chain) is a mint transport failure. Walks the chain because some sites re-raise transport @@ -297,6 +312,13 @@ def classify_redemption_error( "Token was redeemed but could not be credited; do not retry", "cashu_token_consumed", ) + if is_source_mint_connection_error(error): + return ( + "mint_unreachable", + 503, + "The mint that issued this Cashu token is unreachable; the token cannot be redeemed at another mint", + "cashu_source_mint_unreachable", + ) if is_mint_connection_error(error): return ( "mint_unreachable", @@ -375,19 +397,45 @@ async def _redeem_same_mint( that, not the face value, or routstr over-credits the user and its wallet drifts insolvent. """ - await _mint_operation( - lambda: wallet.load_mint(keyset_id=token_obj.keysets[0]), - op_name="redeem_load_mint", - mint_url=token_obj.mint, - ) - wallet.verify_proofs_dleq(token_obj.proofs) - input_fees = wallet.get_fees_for_proofs(token_obj.proofs) - await _mint_operation( - lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), - op_name="redeem_split", - mint_url=token_obj.mint, - retry_timeouts=False, - ) + try: + await _mint_operation( + lambda: wallet.load_mint(keyset_id=token_obj.keysets[0]), + op_name="redeem_load_mint", + mint_url=token_obj.mint, + ) + wallet.verify_proofs_dleq(token_obj.proofs) + input_fees = wallet.get_fees_for_proofs(token_obj.proofs) + await _mint_operation( + lambda: wallet.split( + proofs=token_obj.proofs, amount=0, include_fees=True + ), + op_name="redeem_split", + mint_url=token_obj.mint, + retry_timeouts=False, + ) + except Exception as error: + if is_mint_connection_error(error): + alternatives = [ + mint for mint in settings.cashu_mints if mint != token_obj.mint + ] + logger.warning( + "Same-mint redemption failed", + extra={ + "event": "cashu_same_mint_redemption_failed", + "source_mint": token_obj.mint, + "source_unit": token_obj.unit, + "source_amount": token_obj.amount, + "cross_mint_fallback_available": bool(alternatives), + "destination_candidates": alternatives, + "error": str(error), + "error_type": type(error).__name__, + }, + ) + raise SourceMintConnectionError( + "Issuing Cashu mint is unreachable" + ) from error + raise + return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint @@ -415,18 +463,56 @@ async def recieve_token( "destination_candidates": destinations, }, ) - return await swap_to_primary_mint(token_obj, wallet) + return await swap_to_trusted_mint(token_obj, wallet) - logger.info( - "Cashu same-mint redemption selected", + destinations = [ + mint + for mint in dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) + if mint != token_obj.mint + ] + logger.warning( + "Trying same-mint Cashu redemption", extra={ "event": "cashu_same_mint_redemption", "source_mint": token_obj.mint, "source_unit": token_obj.unit, "source_amount": token_obj.amount, + "cross_mint_fallback_available": bool(destinations), + "destination_candidates": destinations, }, ) - return await _redeem_same_mint(wallet, token_obj) + try: + return await _redeem_same_mint(wallet, token_obj) + except SourceMintConnectionError as same_mint_error: + if not destinations: + raise + logger.warning( + "Same-mint redemption failed; trying cross-mint swap", + extra={ + "event": "cashu_cross_mint_fallback_started", + "source_mint": token_obj.mint, + "source_unit": token_obj.unit, + "source_amount": token_obj.amount, + "destination_candidates": destinations, + "same_mint_error": str(same_mint_error), + }, + ) + try: + return await swap_to_trusted_mint( + token_obj, wallet, force_cross_mint=True + ) + except Exception as swap_error: + logger.error( + "Cross-mint fallback failed", + extra={ + "event": "cashu_cross_mint_fallback_failed", + "source_mint": token_obj.mint, + "destination_candidates": destinations, + "error": str(swap_error), + "error_type": type(swap_error).__name__, + }, + ) + raise async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: @@ -602,7 +688,11 @@ def _melt_insufficient_shortfall(error: Exception) -> int | None: async def _request_mint_with_fallback( - amount: int, *, op_name: str, primary_wallet: Wallet | None = None + amount: int, + *, + op_name: str, + primary_wallet: Wallet | None = None, + excluded_mints: set[str] | None = None, ) -> tuple[Wallet, str, MintQuote]: """Try request_mint on the primary mint, fall back to other trusted mints on transport or rate-limit failure. Returns the wallet, mint_url, and quote. @@ -616,9 +706,13 @@ async def _request_mint_with_fallback( f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " f"Token value is too small after fee deduction or unit conversion." ) - candidates = [settings.primary_mint] + [ - m for m in settings.cashu_mints if m != settings.primary_mint + excluded_mints = excluded_mints or set() + candidates = [ + mint + for mint in [settings.primary_mint, *settings.cashu_mints] + if mint not in excluded_mints ] + candidates = list(dict.fromkeys(candidates)) logger.warning( "Trying trusted destination mints", extra={ @@ -729,6 +823,7 @@ async def _calculate_swap_amount( token_wallet: Wallet, primary_wallet: Wallet | None, proofs: list, + excluded_mints: set[str] | None = None, ) -> int: """ Calculate the amount to mint on the primary mint after accounting for @@ -739,7 +834,7 @@ async def _calculate_swap_amount( else: receive_amount = amount_msat - if token_mint_url == settings.primary_mint: + if token_mint_url == settings.primary_mint and not excluded_mints: logger.info( "swap_to_primary_mint: skipping fee estimation (same mint)", extra={"minted_amount": receive_amount}, @@ -786,6 +881,7 @@ async def _calculate_swap_amount( receive_amount, op_name="swap_fee_est_mint_quote", primary_wallet=primary_wallet, + excluded_mints=excluded_mints, ) stage = "source_fee_quote" dummy_melt_quote = await _mint_operation( @@ -842,14 +938,22 @@ async def _calculate_swap_amount( "event": "cashu_source_mint_unreachable", "source_mint": token_mint_url, "stage": stage, + "fallback_possible": False, + "reason": "cashu_proofs_are_bound_to_the_issuing_mint", }, ) + raise SourceMintConnectionError( + "Issuing Cashu mint is unreachable" + ) from e raise MintConnectionError("Cashu mint is unreachable") from e raise ValueError(f"Failed to estimate fees: {e}") from e -async def swap_to_primary_mint( - token_obj: Token, token_wallet: Wallet +async def swap_to_trusted_mint( + token_obj: Token, + token_wallet: Wallet, + *, + force_cross_mint: bool = False, ) -> tuple[int, str, str]: logger.warning( "Starting Cashu cross-mint swap", @@ -876,7 +980,7 @@ async def swap_to_primary_mint( # If the token is already from the primary mint, we don't need a cross-mint # swap — redeem it same-mint. There's no melt/Lightning fee, but the mint's # NUT-02 input fee still applies; _redeem_same_mint accounts for it. - if token_obj.mint == settings.primary_mint: + if token_obj.mint == settings.primary_mint and not force_cross_mint: logger.info( "swap_to_primary_mint: token already on primary mint, skipping swap", extra={ @@ -888,6 +992,7 @@ async def swap_to_primary_mint( return await _redeem_same_mint(token_wallet, token_obj) primary_wallet: Wallet | None = None + excluded_mints = {token_obj.mint} if force_cross_mint else None minted_amount = await _calculate_swap_amount( amount_msat, @@ -896,6 +1001,7 @@ async def swap_to_primary_mint( token_wallet, primary_wallet, token_obj.proofs, + excluded_mints, ) # The estimate above is non-binding: the mint may demand a higher fee on the @@ -926,7 +1032,10 @@ async def swap_to_primary_mint( f"minted_amount={minted_amount} after fee deduction (attempt {attempt})" ) dest_wallet, dest_mint_url, mint_quote = await _request_mint_with_fallback( - minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet + minted_amount, + op_name="swap_request_mint", + primary_wallet=primary_wallet, + excluded_mints=excluded_mints, ) logger.info( "swap_to_primary_mint: mint quote received", @@ -966,7 +1075,9 @@ async def swap_to_primary_mint( "attempt": attempt, }, ) - raise MintConnectionError("Cashu mint is unreachable") from error + raise SourceMintConnectionError( + "Issuing Cashu mint is unreachable" + ) from error raise input_fees = token_wallet.get_fees_for_proofs(token_obj.proofs) total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees @@ -1046,7 +1157,9 @@ async def swap_to_primary_mint( "attempt": attempt, }, ) - raise MintConnectionError("Cashu mint is unreachable") from e + raise SourceMintConnectionError( + "Issuing Cashu mint is unreachable" + ) from e shortfall = _melt_insufficient_shortfall(e) recomputed = 0 if shortfall is not None: @@ -1184,6 +1297,20 @@ async def swap_to_primary_mint( return int(minted_amount), settings.primary_mint_unit, dest_mint_url +async def swap_to_primary_mint( + token_obj: Token, + token_wallet: Wallet, + *, + force_cross_mint: bool = False, +) -> tuple[int, str, str]: + """Backward-compatible alias for callers using the old function name.""" + return await swap_to_trusted_mint( + token_obj, + token_wallet, + force_cross_mint=force_cross_mint, + ) + + async def credit_balance( cashu_token: str, key: db.ApiKey, session: db.AsyncSession ) -> int: diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 609e2557..0cf78b18 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -534,6 +534,29 @@ async def test_topup_mint_unreachable_returns_503(error: Exception) -> None: assert exc_info.value.detail == "Cashu mint is unreachable" +@pytest.mark.asyncio +async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> None: + from fastapi import HTTPException + + from routstr.wallet import SourceMintConnectionError + + key = _make_api_key(balance=1000) + session = MagicMock() + error = SourceMintConnectionError("Issuing Cashu mint is unreachable") + + with ( + patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)), + patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)), + ): + with pytest.raises(HTTPException) as exc_info: + await topup_wallet_endpoint( + cashu_token="cashuAtoken", key=key, session=session + ) + + assert exc_info.value.status_code == 503 + assert "cannot be redeemed at another mint" in exc_info.value.detail + + @pytest.mark.asyncio async def test_topup_already_spent_still_returns_400() -> None: """Regression: the mint-unreachable short-circuit must not swallow the diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index ff666e65..1a00e630 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -159,6 +159,64 @@ async def test_recieve_token_trusted_mint_deducts_input_fee() -> None: ) +@pytest.mark.asyncio +async def test_primary_mint_failure_falls_back_to_secondary_swap() -> None: + from routstr.core.settings import settings + + source = "http://primary:3338" + destination = "http://secondary:3338" + token = Mock( + mint=source, + unit="sat", + amount=100, + keysets=["keyset1"], + proofs=[Mock(amount=100)], + ) + source_wallet = Mock( + load_mint=AsyncMock(side_effect=httpx.ConnectError("split endpoint down")), + get_fees_for_proofs=Mock(return_value=0), + melt_quote=AsyncMock( + return_value=Mock(quote="melt_quote", amount=90, fee_reserve=10) + ), + melt=AsyncMock(return_value=Mock()), + ) + mint_quote = Mock(quote="mint_quote", request="lnbc1destination") + destination_wallet = Mock( + request_mint=AsyncMock(return_value=mint_quote), + load_proofs=AsyncMock(), + available_balance=Mock(amount=0), + mint=AsyncMock(return_value=Mock()), + keysets=["destination_keyset"], + ) + + async def get_wallet(mint: str, *args: object, **kwargs: object) -> Mock: + return source_wallet if mint == source else destination_wallet + + with ( + patch.object(settings, "primary_mint", source), + patch.object(settings, "primary_mint_unit", "sat"), + patch.object(settings, "cashu_mints", [source, destination]), + patch.object(settings, "mint_max_concurrency", 0), + patch.object(settings, "mint_operation_timeout_seconds", 0), + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch("routstr.wallet.get_wallet", AsyncMock(side_effect=get_wallet)), + patch("routstr.wallet.logger.warning") as warning, + ): + amount, unit, mint = await recieve_token("cashuAtoken") + + assert (amount, unit, mint) == (90, "sat", destination) + source_wallet.melt.assert_awaited_once() + destination_wallet.mint.assert_awaited_once() + events = { + call.kwargs["extra"].get("event") + for call in warning.call_args_list + if "extra" in call.kwargs + } + assert "cashu_cross_mint_fallback_started" in events + assert "cashu_destination_selected" in events + assert "cashu_swap_completed" in events + + @pytest.mark.asyncio async def test_send_token() -> None: mock_wallet = Mock() @@ -393,7 +451,7 @@ async def test_recieve_token_untrusted_mint() -> None: mock_wallet.load_proofs = AsyncMock() with patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet): with patch( - "routstr.wallet.swap_to_primary_mint", + "routstr.wallet.swap_to_trusted_mint", return_value=(900, "sat", "http://mint:3338"), ): amount, unit, mint = await recieve_token("test_token") From 39970d8bee1eae875b0bc58bebad824c4b6580bf Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 14 Jul 2026 01:23:14 +0200 Subject: [PATCH 14/46] clean up --- routstr/wallet.py | 91 ++++++--------------------------------- tests/unit/test_wallet.py | 48 +++++++-------------- 2 files changed, 29 insertions(+), 110 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index a898ab7c..54a57978 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -415,18 +415,15 @@ async def _redeem_same_mint( ) except Exception as error: if is_mint_connection_error(error): - alternatives = [ - mint for mint in settings.cashu_mints if mint != token_obj.mint - ] logger.warning( - "Same-mint redemption failed", + "Same-mint redemption failed; client must use a different token", extra={ "event": "cashu_same_mint_redemption_failed", "source_mint": token_obj.mint, "source_unit": token_obj.unit, "source_amount": token_obj.amount, - "cross_mint_fallback_available": bool(alternatives), - "destination_candidates": alternatives, + "cross_mint_fallback_attempted": False, + "action": "retry_with_token_from_another_mint", "error": str(error), "error_type": type(error).__name__, }, @@ -465,11 +462,6 @@ async def recieve_token( ) return await swap_to_trusted_mint(token_obj, wallet) - destinations = [ - mint - for mint in dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) - if mint != token_obj.mint - ] logger.warning( "Trying same-mint Cashu redemption", extra={ @@ -477,42 +469,10 @@ async def recieve_token( "source_mint": token_obj.mint, "source_unit": token_obj.unit, "source_amount": token_obj.amount, - "cross_mint_fallback_available": bool(destinations), - "destination_candidates": destinations, + "cross_mint_fallback_on_connection_failure": False, }, ) - try: - return await _redeem_same_mint(wallet, token_obj) - except SourceMintConnectionError as same_mint_error: - if not destinations: - raise - logger.warning( - "Same-mint redemption failed; trying cross-mint swap", - extra={ - "event": "cashu_cross_mint_fallback_started", - "source_mint": token_obj.mint, - "source_unit": token_obj.unit, - "source_amount": token_obj.amount, - "destination_candidates": destinations, - "same_mint_error": str(same_mint_error), - }, - ) - try: - return await swap_to_trusted_mint( - token_obj, wallet, force_cross_mint=True - ) - except Exception as swap_error: - logger.error( - "Cross-mint fallback failed", - extra={ - "event": "cashu_cross_mint_fallback_failed", - "source_mint": token_obj.mint, - "destination_candidates": destinations, - "error": str(swap_error), - "error_type": type(swap_error).__name__, - }, - ) - raise + return await _redeem_same_mint(wallet, token_obj) async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: @@ -688,11 +648,7 @@ def _melt_insufficient_shortfall(error: Exception) -> int | None: async def _request_mint_with_fallback( - amount: int, - *, - op_name: str, - primary_wallet: Wallet | None = None, - excluded_mints: set[str] | None = None, + amount: int, *, op_name: str, primary_wallet: Wallet | None = None ) -> tuple[Wallet, str, MintQuote]: """Try request_mint on the primary mint, fall back to other trusted mints on transport or rate-limit failure. Returns the wallet, mint_url, and quote. @@ -706,13 +662,9 @@ async def _request_mint_with_fallback( f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " f"Token value is too small after fee deduction or unit conversion." ) - excluded_mints = excluded_mints or set() - candidates = [ - mint - for mint in [settings.primary_mint, *settings.cashu_mints] - if mint not in excluded_mints - ] - candidates = list(dict.fromkeys(candidates)) + candidates = list( + dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) + ) logger.warning( "Trying trusted destination mints", extra={ @@ -823,7 +775,6 @@ async def _calculate_swap_amount( token_wallet: Wallet, primary_wallet: Wallet | None, proofs: list, - excluded_mints: set[str] | None = None, ) -> int: """ Calculate the amount to mint on the primary mint after accounting for @@ -834,7 +785,7 @@ async def _calculate_swap_amount( else: receive_amount = amount_msat - if token_mint_url == settings.primary_mint and not excluded_mints: + if token_mint_url == settings.primary_mint: logger.info( "swap_to_primary_mint: skipping fee estimation (same mint)", extra={"minted_amount": receive_amount}, @@ -881,7 +832,6 @@ async def _calculate_swap_amount( receive_amount, op_name="swap_fee_est_mint_quote", primary_wallet=primary_wallet, - excluded_mints=excluded_mints, ) stage = "source_fee_quote" dummy_melt_quote = await _mint_operation( @@ -950,10 +900,7 @@ async def _calculate_swap_amount( async def swap_to_trusted_mint( - token_obj: Token, - token_wallet: Wallet, - *, - force_cross_mint: bool = False, + token_obj: Token, token_wallet: Wallet ) -> tuple[int, str, str]: logger.warning( "Starting Cashu cross-mint swap", @@ -980,7 +927,7 @@ async def swap_to_trusted_mint( # If the token is already from the primary mint, we don't need a cross-mint # swap — redeem it same-mint. There's no melt/Lightning fee, but the mint's # NUT-02 input fee still applies; _redeem_same_mint accounts for it. - if token_obj.mint == settings.primary_mint and not force_cross_mint: + if token_obj.mint == settings.primary_mint: logger.info( "swap_to_primary_mint: token already on primary mint, skipping swap", extra={ @@ -992,7 +939,6 @@ async def swap_to_trusted_mint( return await _redeem_same_mint(token_wallet, token_obj) primary_wallet: Wallet | None = None - excluded_mints = {token_obj.mint} if force_cross_mint else None minted_amount = await _calculate_swap_amount( amount_msat, @@ -1001,7 +947,6 @@ async def swap_to_trusted_mint( token_wallet, primary_wallet, token_obj.proofs, - excluded_mints, ) # The estimate above is non-binding: the mint may demand a higher fee on the @@ -1035,7 +980,6 @@ async def swap_to_trusted_mint( minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet, - excluded_mints=excluded_mints, ) logger.info( "swap_to_primary_mint: mint quote received", @@ -1298,17 +1242,10 @@ async def swap_to_trusted_mint( async def swap_to_primary_mint( - token_obj: Token, - token_wallet: Wallet, - *, - force_cross_mint: bool = False, + token_obj: Token, token_wallet: Wallet ) -> tuple[int, str, str]: """Backward-compatible alias for callers using the old function name.""" - return await swap_to_trusted_mint( - token_obj, - token_wallet, - force_cross_mint=force_cross_mint, - ) + return await swap_to_trusted_mint(token_obj, token_wallet) async def credit_balance( diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 1a00e630..d851cdaa 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -160,8 +160,9 @@ async def test_recieve_token_trusted_mint_deducts_input_fee() -> None: @pytest.mark.asyncio -async def test_primary_mint_failure_falls_back_to_secondary_swap() -> None: +async def test_primary_mint_failure_does_not_try_another_mint() -> None: from routstr.core.settings import settings + from routstr.wallet import SourceMintConnectionError source = "http://primary:3338" destination = "http://secondary:3338" @@ -173,48 +174,29 @@ async def test_primary_mint_failure_falls_back_to_secondary_swap() -> None: proofs=[Mock(amount=100)], ) source_wallet = Mock( - load_mint=AsyncMock(side_effect=httpx.ConnectError("split endpoint down")), - get_fees_for_proofs=Mock(return_value=0), - melt_quote=AsyncMock( - return_value=Mock(quote="melt_quote", amount=90, fee_reserve=10) - ), - melt=AsyncMock(return_value=Mock()), + load_mint=AsyncMock(side_effect=httpx.ConnectError("mint unavailable")) ) - mint_quote = Mock(quote="mint_quote", request="lnbc1destination") - destination_wallet = Mock( - request_mint=AsyncMock(return_value=mint_quote), - load_proofs=AsyncMock(), - available_balance=Mock(amount=0), - mint=AsyncMock(return_value=Mock()), - keysets=["destination_keyset"], - ) - - async def get_wallet(mint: str, *args: object, **kwargs: object) -> Mock: - return source_wallet if mint == source else destination_wallet + get_wallet = AsyncMock(return_value=source_wallet) with ( patch.object(settings, "primary_mint", source), - patch.object(settings, "primary_mint_unit", "sat"), patch.object(settings, "cashu_mints", [source, destination]), - patch.object(settings, "mint_max_concurrency", 0), - patch.object(settings, "mint_operation_timeout_seconds", 0), patch("routstr.wallet.deserialize_token_from_string", return_value=token), - patch("routstr.wallet.get_wallet", AsyncMock(side_effect=get_wallet)), + patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.logger.warning") as warning, ): - amount, unit, mint = await recieve_token("cashuAtoken") + with pytest.raises(SourceMintConnectionError): + await recieve_token("cashuAtoken") - assert (amount, unit, mint) == (90, "sat", destination) - source_wallet.melt.assert_awaited_once() - destination_wallet.mint.assert_awaited_once() - events = { - call.kwargs["extra"].get("event") + get_wallet.assert_awaited_once_with(source, "sat", load=False) + failure = next( + call.kwargs["extra"] for call in warning.call_args_list - if "extra" in call.kwargs - } - assert "cashu_cross_mint_fallback_started" in events - assert "cashu_destination_selected" in events - assert "cashu_swap_completed" in events + if call.kwargs.get("extra", {}).get("event") + == "cashu_same_mint_redemption_failed" + ) + assert failure["cross_mint_fallback_attempted"] is False + assert failure["action"] == "retry_with_token_from_another_mint" @pytest.mark.asyncio From 93ab1d927bcd7f60b5638edc6ee68e0a2e8b52be Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 14 Jul 2026 01:35:54 +0200 Subject: [PATCH 15/46] mint cooldown --- routstr/wallet.py | 17 +++++++++++++---- tests/unit/test_fetch_all_balances.py | 27 +++++++++++++++++++++++++++ tests/unit/test_wallet.py | 17 +++++++++++++++++ 3 files changed, 57 insertions(+), 4 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index 54a57978..bd59de91 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -319,7 +319,7 @@ def classify_redemption_error( "The mint that issued this Cashu token is unreachable; the token cannot be redeemed at another mint", "cashu_source_mint_unreachable", ) - if is_mint_connection_error(error): + if _is_mint_rate_limited(error) or is_mint_connection_error(error): return ( "mint_unreachable", 503, @@ -1456,7 +1456,7 @@ class BalanceDetail(TypedDict, total=False): _BALANCE_FETCH_RETRY_SECONDS = 60.0 _balance_fetch_failures: dict[tuple[str, str], tuple[float, str]] = {} -_balance_fetch_locks: dict[tuple[str, str], asyncio.Lock] = {} +_balance_fetch_locks: dict[str, asyncio.Lock] = {} def _balance_error(mint_url: str, unit: str, error: str) -> BalanceDetail: @@ -1489,7 +1489,7 @@ async def fetch_all_balances( session: db.AsyncSession, mint_url: str, unit: str ) -> BalanceDetail: key = (mint_url, unit) - lock = _balance_fetch_locks.setdefault(key, asyncio.Lock()) + lock = _balance_fetch_locks.setdefault(mint_url, asyncio.Lock()) async with lock: now = time.monotonic() failure = _balance_fetch_failures.get(key) @@ -1498,7 +1498,7 @@ async def fetch_all_balances( cooldown = _mint_cooldown_remaining(mint_url) if cooldown > 0: - error = "Mint is cooling down after a rate limit" + error = "Mint cooldown is active" _balance_fetch_failures[key] = (now + cooldown, error) return _balance_error(mint_url, unit, error) @@ -1516,6 +1516,12 @@ async def fetch_all_balances( session, mint_url, unit ) except Exception as error: + connection_failure = is_mint_connection_error(error) + rate_limited = _is_mint_rate_limited(error) + if connection_failure or rate_limited: + _MintRateGuard.get(mint_url).apply_cooldown( + _BALANCE_FETCH_RETRY_SECONDS + ) retry_delay = max( _BALANCE_FETCH_RETRY_SECONDS, _mint_cooldown_remaining(mint_url), @@ -1528,6 +1534,9 @@ async def fetch_all_balances( "mint_url": mint_url, "unit": unit, "error": str(error), + "connection_failure": connection_failure, + "rate_limited": rate_limited, + "mint_cooldown_applied": connection_failure or rate_limited, "retry_seconds": round(retry_delay, 2), }, ) diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index 20cbe312..a088c0b7 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -14,9 +14,11 @@ def clear_balance_fetch_state() -> Generator[None, None, None]: wallet._balance_fetch_failures.clear() wallet._balance_fetch_locks.clear() + wallet._MintRateGuard._guards.clear() yield wallet._balance_fetch_failures.clear() wallet._balance_fetch_locks.clear() + wallet._MintRateGuard._guards.clear() @asynccontextmanager @@ -99,6 +101,31 @@ async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: assert get_wallet.await_count == 2 +@pytest.mark.asyncio +async def test_balance_failure_applies_mint_cooldown_to_other_units() -> None: + from routstr.core.settings import settings + from routstr.wallet import _mint_cooldown_remaining + + mint = "http://mint:3338" + get_wallet = AsyncMock(side_effect=httpx.ConnectError("mint unavailable")) + with ( + patch.object(settings, "cashu_mints", [mint]), + patch.object(settings, "primary_mint", mint), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.db.create_session", _fake_session), + patch("routstr.wallet.time.monotonic", return_value=10), + patch("routstr.wallet.logger.warning") as warning, + ): + details, *_ = await fetch_all_balances(units=["sat", "msat"]) + cooldown = _mint_cooldown_remaining(mint) + + assert get_wallet.await_count == 1 + assert warning.call_count == 1 + assert cooldown == 60 + assert details[0]["error"] == "mint unavailable" + assert details[1]["error"] == "Mint cooldown is active" + + @pytest.mark.asyncio async def test_fetch_all_balances_no_duplicate_primary_mint() -> None: """primary_mint already in cashu_mints is not inspected twice.""" diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index d851cdaa..a192eeac 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1229,6 +1229,23 @@ def _chain(outer: BaseException, cause: BaseException) -> BaseException: return outer +def test_rate_limited_mint_is_classified_as_unreachable() -> None: + from routstr.wallet import classify_redemption_error + + request = httpx.Request("POST", "http://mint:3338/v1/swap") + response = httpx.Response(429, request=request) + error = httpx.HTTPStatusError( + "rate limited", request=request, response=response + ) + + assert classify_redemption_error(error) == ( + "mint_unreachable", + 503, + "Cashu mint is unreachable", + "cashu_mint_unreachable", + ) + + @pytest.mark.parametrize( "error", [ From cc2a96e2ef8d61f79bb5f8f0ab92a964714555c1 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 14 Jul 2026 21:59:43 +0200 Subject: [PATCH 16/46] make trusted mint available for lightning topup --- routstr/lightning.py | 16 ++++++++++++---- .../integration/test_lightning_invoice_rip08.py | 8 +++++--- 2 files changed, 17 insertions(+), 7 deletions(-) diff --git a/routstr/lightning.py b/routstr/lightning.py index fafffef9..c4bc2b20 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -71,6 +71,14 @@ class InvoiceRecoverRequest(BaseModel): bolt11: str = Field(description="BOLT11 invoice string") +def _trusted_mint_candidates() -> list[str]: + return [ + mint + for mint in dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) + if mint + ] + + async def _request_mint_with_fallback( amount_sats: int, *, @@ -167,13 +175,13 @@ async def create_invoice( try: description = f"Routstr {request.purpose} {request.amount_sats} sats" - # An API key is backed by one mint. A top-up must use that same mint; - # falling back to another would create mixed-mint collateral that the - # current single refund_mint_url field cannot account for or refund. allowed_mints = None if request.purpose == "topup": assert topup_api_key is not None - allowed_mints = [topup_api_key.refund_mint_url or settings.primary_mint] + # Top-ups are not pinned to the key's previous/backing mint. Use any + # currently available trusted mint so rate limits/cooldowns on one + # mint do not block Lightning top-ups. + allowed_mints = _trusted_mint_candidates() bolt11, payment_hash, mint_url = await generate_lightning_invoice( request.amount_sats, description, allowed_mints=allowed_mints ) diff --git a/tests/integration/test_lightning_invoice_rip08.py b/tests/integration/test_lightning_invoice_rip08.py index 766d0176..99634f12 100644 --- a/tests/integration/test_lightning_invoice_rip08.py +++ b/tests/integration/test_lightning_invoice_rip08.py @@ -16,6 +16,7 @@ from httpx import AsyncClient from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.db import ApiKey +from routstr.core.settings import settings RIP08_PATH = "/lightning/invoice" LEGACY_PATH = "/v1/balance/lightning/invoice" @@ -101,9 +102,10 @@ async def test_topup_with_authorization_header( body = resp.json() assert body["amount_sats"] == 500 assert body["bolt11"].startswith("lnbc") - assert patch_invoice_generation.call_args.kwargs["allowed_mints"] == [ - "http://localhost:3338" - ] + expected_mints = [settings.primary_mint, *settings.cashu_mints] + allowed_mints = patch_invoice_generation.call_args.kwargs["allowed_mints"] + assert allowed_mints == list(dict.fromkeys(mint for mint in expected_mints if mint)) + assert allowed_mints[0] == settings.primary_mint @pytest.mark.integration From 09e1c7bf2d73582f49ade48397989e4b655be357 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 15 Jul 2026 00:10:17 +0200 Subject: [PATCH 17/46] better cooldown --- routstr/wallet.py | 68 ++++++++++++++++++++++++++++++++++----- tests/unit/test_wallet.py | 29 +++++++++++++++++ 2 files changed, 89 insertions(+), 8 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index bd59de91..e3384bfe 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -89,30 +89,82 @@ class _MintRateGuard: asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None ) self._cooldown_until = 0.0 + self._needs_probe = False + self._probe_lock = asyncio.Lock() def apply_cooldown(self, delay: float) -> None: self._cooldown_until = max( self._cooldown_until, time.monotonic() + max(0.0, delay) ) + self._needs_probe = True def cooldown_remaining(self) -> float: return max(0.0, self._cooldown_until - time.monotonic()) - async def _run_after_cooldown(self, factory: Callable[[], Awaitable[Any]]) -> Any: - wait = self.cooldown_remaining() - if wait > 0: + async def _wait_for_cooldown(self) -> None: + while True: + deadline = self._cooldown_until + wait = max(0.0, deadline - time.monotonic()) + if wait <= 0: + return logger.debug( "Mint rate guard: cooling down", extra={"mint_url": self._mint_url, "wait_seconds": round(wait, 2)}, ) await asyncio.sleep(wait) - return await factory() + if self._cooldown_until <= deadline: + return + + async def _run_probe(self, factory: Callable[[], Awaitable[Any]]) -> Any: + await self._wait_for_cooldown() + logger.warning( + "Mint cooldown ended; sending one probe request", + extra={"event": "mint_cooldown_probe_started", "mint_url": self._mint_url}, + ) + try: + result = await factory() + except Exception as error: + # Keep queued callers behind the probe while _mint_operation applies + # the precise Retry-After/backoff from this failure. + self.apply_cooldown(1.0) + logger.warning( + "Mint cooldown probe failed", + extra={ + "event": "mint_cooldown_probe_failed", + "mint_url": self._mint_url, + "error": str(error), + "error_type": type(error).__name__, + }, + ) + raise + + self._needs_probe = False + self._cooldown_until = 0.0 + logger.warning( + "Mint cooldown probe succeeded; restoring normal concurrency", + extra={ + "event": "mint_cooldown_probe_succeeded", + "mint_url": self._mint_url, + }, + ) + return result async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any: - if self._semaphore is None: - return await self._run_after_cooldown(factory) - async with self._semaphore: - return await self._run_after_cooldown(factory) + while True: + if self._needs_probe or self.cooldown_remaining() > 0: + async with self._probe_lock: + if self.cooldown_remaining() > 0: + self._needs_probe = True + if self._needs_probe: + return await self._run_probe(factory) + continue + + if self._semaphore is None: + return await factory() + async with self._semaphore: + if self._needs_probe: + continue + return await factory() def _mint_cooldown_remaining(mint_url: str) -> float: diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index a192eeac..fd3e3011 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1511,6 +1511,35 @@ async def test_mint_rate_guard_waits_for_adaptive_cooldown() -> None: operation.assert_awaited_once() +@pytest.mark.asyncio +async def test_mint_rate_guard_allows_one_probe_after_cooldown() -> None: + from routstr.wallet import _MintRateGuard + + guard = _MintRateGuard("http://mint:3338", 4) + guard.apply_cooldown(0) + probe_started = asyncio.Event() + release_probe = asyncio.Event() + calls = 0 + + async def operation() -> int: + nonlocal calls + calls += 1 + if calls == 1: + probe_started.set() + await release_probe.wait() + return calls + + tasks = [asyncio.create_task(guard.run(operation)) for _ in range(5)] + await probe_started.wait() + await asyncio.sleep(0) + assert calls == 1 + + release_probe.set() + await asyncio.gather(*tasks) + assert calls == 5 + assert guard._needs_probe is False + + def test_mint_rate_guard_rebuilds_when_setting_changes() -> None: from routstr.core.settings import settings from routstr.wallet import _MintRateGuard From 8b942f3c149d846e537e7c5f12893663f8769d3b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 15 Jul 2026 01:00:32 +0200 Subject: [PATCH 18/46] show mint status correclty --- routstr/wallet.py | 183 +++++++++++++++++----- tests/unit/test_fetch_all_balances.py | 112 ++++++++++++- tests/unit/test_wallet.py | 4 +- ui/components/detailed-wallet-balance.tsx | 31 +++- ui/lib/api/services/wallet.ts | 2 + 5 files changed, 284 insertions(+), 48 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index e3384bfe..8828f04c 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -89,18 +89,26 @@ class _MintRateGuard: asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None ) self._cooldown_until = 0.0 + self._cooldown_reason: str | None = None self._needs_probe = False self._probe_lock = asyncio.Lock() - def apply_cooldown(self, delay: float) -> None: - self._cooldown_until = max( - self._cooldown_until, time.monotonic() + max(0.0, delay) - ) + def apply_cooldown(self, delay: float, *, reason: str | None = None) -> None: + deadline = time.monotonic() + max(0.0, delay) + if deadline >= self._cooldown_until: + self._cooldown_until = deadline + if reason is not None: + self._cooldown_reason = reason + elif self._cooldown_reason is None and reason is not None: + self._cooldown_reason = reason self._needs_probe = True def cooldown_remaining(self) -> float: return max(0.0, self._cooldown_until - time.monotonic()) + def cooldown_reason(self) -> str | None: + return self._cooldown_reason if self.cooldown_remaining() > 0 else None + async def _wait_for_cooldown(self) -> None: while True: deadline = self._cooldown_until @@ -140,6 +148,7 @@ class _MintRateGuard: self._needs_probe = False self._cooldown_until = 0.0 + self._cooldown_reason = None logger.warning( "Mint cooldown probe succeeded; restoring normal concurrency", extra={ @@ -171,6 +180,10 @@ def _mint_cooldown_remaining(mint_url: str) -> float: return _MintRateGuard.get(mint_url).cooldown_remaining() +def _mint_cooldown_reason(mint_url: str) -> str | None: + return _MintRateGuard.get(mint_url).cooldown_reason() + + def _is_mint_rate_limited(error: BaseException) -> bool: """True if the mint returned a 429 or rate-limit indication.""" current: BaseException | None = error @@ -248,7 +261,7 @@ async def _mint_operation( if retry_after is not None: backoff = max(retry_after, backoff) if guard is not None: - guard.apply_cooldown(backoff) + guard.apply_cooldown(backoff, reason="rate_limited") # When the caller has a fallback strategy (trusted-mint # list), re-raise immediately so the caller can try the next @@ -458,9 +471,7 @@ async def _redeem_same_mint( wallet.verify_proofs_dleq(token_obj.proofs) input_fees = wallet.get_fees_for_proofs(token_obj.proofs) await _mint_operation( - lambda: wallet.split( - proofs=token_obj.proofs, amount=0, include_fees=True - ), + lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), op_name="redeem_split", mint_url=token_obj.mint, retry_timeouts=False, @@ -714,9 +725,7 @@ async def _request_mint_with_fallback( f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " f"Token value is too small after fee deduction or unit conversion." ) - candidates = list( - dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) - ) + candidates = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) logger.warning( "Trying trusted destination mints", extra={ @@ -788,7 +797,7 @@ async def _request_mint_with_fallback( raise if connection_failure: _MintRateGuard.get(mint_url).apply_cooldown( - _MINT_TRANSPORT_COOLDOWN_SECONDS + _MINT_TRANSPORT_COOLDOWN_SECONDS, reason="unreachable" ) logger.warning( "Destination mint failed", @@ -1504,22 +1513,68 @@ class BalanceDetail(TypedDict, total=False): user_balance: int owner_balance: int error: str + error_code: str + retry_after_seconds: float _BALANCE_FETCH_RETRY_SECONDS = 60.0 -_balance_fetch_failures: dict[tuple[str, str], tuple[float, str]] = {} +_MINT_UNITS_CACHE_SECONDS = 300.0 +_balance_fetch_failures: dict[tuple[str, str], tuple[float, str, str]] = {} _balance_fetch_locks: dict[str, asyncio.Lock] = {} +_mint_supported_units: dict[str, tuple[float, list[str]]] = {} -def _balance_error(mint_url: str, unit: str, error: str) -> BalanceDetail: - return { +async def _get_supported_mint_units(mint_url: str) -> list[str]: + now = time.monotonic() + cached = _mint_supported_units.get(mint_url) + if cached is not None and now < cached[0]: + return cached[1] + + wallet = await get_wallet(mint_url, settings.primary_mint_unit, load=False) + keysets = await _mint_operation( + lambda: wallet._get_keysets(), + op_name="get_mint_keysets", + mint_url=mint_url, + retry_on_rate_limit=False, + ) + units = list( + dict.fromkeys( + keyset.unit.name for keyset in keysets if keyset.active and keyset.unit.name + ) + ) + if not units: + units = [settings.primary_mint_unit] + elif settings.primary_mint_unit in units: + units.remove(settings.primary_mint_unit) + units.insert(0, settings.primary_mint_unit) + + _mint_supported_units[mint_url] = ( + time.monotonic() + _MINT_UNITS_CACHE_SECONDS, + units, + ) + return units + + +def _balance_error( + mint_url: str, + unit: str, + error: str, + *, + error_code: str, + retry_after_seconds: float | None = None, +) -> BalanceDetail: + detail: BalanceDetail = { "mint_url": mint_url, "unit": unit, "wallet_balance": 0, "user_balance": 0, "owner_balance": 0, "error": error, + "error_code": error_code, } + if retry_after_seconds is not None: + detail["retry_after_seconds"] = round(max(0.0, retry_after_seconds), 2) + return detail async def fetch_all_balances( @@ -1534,8 +1589,6 @@ async def fetch_all_balances( - Total user balance in sats - Owner balance in sats (wallet - user) """ - if units is None: - units = ["sat", "msat"] async def fetch_balance( session: db.AsyncSession, mint_url: str, unit: str @@ -1546,18 +1599,36 @@ async def fetch_all_balances( now = time.monotonic() failure = _balance_fetch_failures.get(key) if failure is not None and now < failure[0]: - return _balance_error(mint_url, unit, failure[1]) + return _balance_error( + mint_url, + unit, + failure[1], + error_code=failure[2], + retry_after_seconds=failure[0] - now, + ) cooldown = _mint_cooldown_remaining(mint_url) if cooldown > 0: - error = "Mint cooldown is active" - _balance_fetch_failures[key] = (now + cooldown, error) - return _balance_error(mint_url, unit, error) + error_code = _mint_cooldown_reason(mint_url) or "cooldown" + error = { + "rate_limited": "Mint is rate limited", + "unreachable": "Mint is unreachable", + }.get(error_code, "Mint cooldown is active") + _balance_fetch_failures[key] = ( + now + cooldown, + error, + error_code, + ) + return _balance_error( + mint_url, + unit, + error, + error_code=error_code, + retry_after_seconds=cooldown, + ) try: - wallet = await get_wallet( - mint_url, unit, retry_on_rate_limit=False - ) + wallet = await get_wallet(mint_url, unit, retry_on_rate_limit=False) proofs = get_proofs_per_mint_and_unit( wallet, mint_url, unit, not_reserved=True ) @@ -1570,16 +1641,23 @@ async def fetch_all_balances( except Exception as error: connection_failure = is_mint_connection_error(error) rate_limited = _is_mint_rate_limited(error) + error_code = ( + "rate_limited" + if rate_limited + else "unreachable" + if connection_failure + else "mint_error" + ) if connection_failure or rate_limited: _MintRateGuard.get(mint_url).apply_cooldown( - _BALANCE_FETCH_RETRY_SECONDS + _BALANCE_FETCH_RETRY_SECONDS, reason=error_code ) retry_delay = max( _BALANCE_FETCH_RETRY_SECONDS, _mint_cooldown_remaining(mint_url), ) retry_at = time.monotonic() + retry_delay - _balance_fetch_failures[key] = (retry_at, str(error)) + _balance_fetch_failures[key] = (retry_at, str(error), error_code) logger.warning( "Unable to refresh mint balance", extra={ @@ -1592,7 +1670,13 @@ async def fetch_all_balances( "retry_seconds": round(retry_delay, 2), }, ) - return _balance_error(mint_url, unit, str(error)) + return _balance_error( + mint_url, + unit, + str(error), + error_code=error_code, + retry_after_seconds=retry_delay, + ) _balance_fetch_failures.pop(key, None) if unit == "sat": @@ -1616,16 +1700,45 @@ async def fetch_all_balances( if settings.primary_mint and settings.primary_mint not in mint_urls: mint_urls.append(settings.primary_mint) - # Create tasks for all mint/unit combinations async with db.create_session() as session: - tasks = [ - fetch_balance(session, mint_url, unit) - for mint_url in mint_urls - for unit in units - ] - # Run all tasks concurrently - balance_details = list(await asyncio.gather(*tasks)) + async def fetch_mint_balances(mint_url: str) -> list[BalanceDetail]: + mint_units = units + if mint_units is None: + try: + mint_units = await _get_supported_mint_units(mint_url) + except Exception as error: + connection_failure = is_mint_connection_error(error) + rate_limited = _is_mint_rate_limited(error) + if connection_failure: + _MintRateGuard.get(mint_url).apply_cooldown( + _BALANCE_FETCH_RETRY_SECONDS, reason="unreachable" + ) + # _mint_operation already records rate-limit cooldowns. + # Fetching the configured unit turns a known cooldown into + # a structured error without another mint request. + mint_units = [settings.primary_mint_unit] + if not connection_failure and not rate_limited: + logger.warning( + "Unable to discover mint units", + extra={ + "mint_url": mint_url, + "error": str(error), + "error_type": type(error).__name__, + }, + ) + return list( + await asyncio.gather( + *(fetch_balance(session, mint_url, unit) for unit in mint_units) + ) + ) + + grouped_details = await asyncio.gather( + *(fetch_mint_balances(mint_url) for mint_url in mint_urls) + ) + balance_details = [ + detail for mint_details in grouped_details for detail in mint_details + ] # Calculate totals total_wallet_balance_sats = 0 diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index a088c0b7..ce640e9d 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -14,10 +14,12 @@ def clear_balance_fetch_state() -> Generator[None, None, None]: wallet._balance_fetch_failures.clear() wallet._balance_fetch_locks.clear() + wallet._mint_supported_units.clear() wallet._MintRateGuard._guards.clear() yield wallet._balance_fetch_failures.clear() wallet._balance_fetch_locks.clear() + wallet._mint_supported_units.clear() wallet._MintRateGuard._guards.clear() @@ -51,8 +53,9 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None: """With empty cashu_mints, balances are still fetched for primary_mint.""" from routstr.core.settings import settings - with patch.object(settings, "cashu_mints", []), patch.object( - settings, "primary_mint", "http://primary:3338" + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), ): for p in _patches(proof_amount=1000): p.start() @@ -67,6 +70,78 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None: assert total_wallet == 1000 +@pytest.mark.asyncio +async def test_fetch_all_balances_uses_units_advertised_by_mint() -> None: + from routstr.core.settings import settings + + with ( + patch.object(settings, "cashu_mints", ["http://mint:3338"]), + patch.object(settings, "primary_mint", "http://mint:3338"), + patch( + "routstr.wallet._get_supported_mint_units", + AsyncMock(return_value=["sat"]), + ) as supported_units, + ): + for p in _patches(proof_amount=1000): + p.start() + try: + details, *_ = await fetch_all_balances() + finally: + patch.stopall() + + supported_units.assert_awaited_once_with("http://mint:3338") + assert [detail["unit"] for detail in details] == ["sat"] + + +@pytest.mark.asyncio +async def test_unit_discovery_failure_returns_structured_balance_error() -> None: + from routstr.core.settings import settings + + get_wallet = AsyncMock() + with ( + patch.object(settings, "cashu_mints", ["http://mint:3338"]), + patch.object(settings, "primary_mint", "http://mint:3338"), + patch( + "routstr.wallet._get_supported_mint_units", + AsyncMock(side_effect=httpx.ConnectError("mint unavailable")), + ), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.db.create_session", _fake_session), + ): + details, *_ = await fetch_all_balances() + + assert details[0]["unit"] == settings.primary_mint_unit + assert details[0]["error_code"] == "unreachable" + assert details[0]["retry_after_seconds"] > 0 + get_wallet.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_supported_mint_units_come_from_active_keysets() -> None: + from routstr.core.settings import settings + from routstr.wallet import _get_supported_mint_units + + sat = MagicMock(active=True) + sat.unit.name = "sat" + msat = MagicMock(active=False) + msat.unit.name = "msat" + usd = MagicMock(active=True) + usd.unit.name = "usd" + wallet = MagicMock() + wallet._get_keysets = AsyncMock(return_value=[usd, msat, sat]) + + with ( + patch.object(settings, "primary_mint_unit", "sat"), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), + ): + units = await _get_supported_mint_units("http://mint:3338") + cached_units = await _get_supported_mint_units("http://mint:3338") + + assert units == ["sat", "usd"] + assert cached_units == units + wallet._get_keysets.assert_awaited_once() + + @pytest.mark.asyncio async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: from routstr.core.settings import settings @@ -84,7 +159,10 @@ async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: second = await fetch_all_balances(units=["sat"]) assert first[0][0]["error"] == "mint unavailable" + assert first[0][0]["error_code"] == "unreachable" + assert first[0][0]["retry_after_seconds"] == 60 assert second[0][0]["error"] == "mint unavailable" + assert second[0][0]["error_code"] == "unreachable" assert get_wallet.await_count == 1 warning.assert_called_once() @@ -101,6 +179,25 @@ async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: assert get_wallet.await_count == 2 +@pytest.mark.asyncio +async def test_fetch_all_balances_reports_rate_limit_status() -> None: + from routstr.core.settings import settings + + request = httpx.Request("GET", "http://mint:3338/v1/keysets") + response = httpx.Response(429, request=request, headers={"Retry-After": "45"}) + error = httpx.HTTPStatusError("rate limited", request=request, response=response) + with ( + patch.object(settings, "cashu_mints", ["http://mint:3338"]), + patch.object(settings, "primary_mint", "http://mint:3338"), + patch("routstr.wallet.get_wallet", AsyncMock(side_effect=error)), + patch("routstr.wallet.db.create_session", _fake_session), + ): + details, *_ = await fetch_all_balances(units=["sat"]) + + assert details[0]["error_code"] == "rate_limited" + assert details[0]["retry_after_seconds"] == 60 + + @pytest.mark.asyncio async def test_balance_failure_applies_mint_cooldown_to_other_units() -> None: from routstr.core.settings import settings @@ -123,7 +220,9 @@ async def test_balance_failure_applies_mint_cooldown_to_other_units() -> None: assert warning.call_count == 1 assert cooldown == 60 assert details[0]["error"] == "mint unavailable" - assert details[1]["error"] == "Mint cooldown is active" + assert details[0]["error_code"] == "unreachable" + assert details[1]["error"] == "Mint is unreachable" + assert details[1]["error_code"] == "unreachable" @pytest.mark.asyncio @@ -131,9 +230,10 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None: """primary_mint already in cashu_mints is not inspected twice.""" from routstr.core.settings import settings - with patch.object( - settings, "cashu_mints", ["http://primary:3338"] - ), patch.object(settings, "primary_mint", "http://primary:3338"): + with ( + patch.object(settings, "cashu_mints", ["http://primary:3338"]), + patch.object(settings, "primary_mint", "http://primary:3338"), + ): for p in _patches(proof_amount=1000): p.start() try: diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index fd3e3011..39d87899 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1234,9 +1234,7 @@ def test_rate_limited_mint_is_classified_as_unreachable() -> None: request = httpx.Request("POST", "http://mint:3338/v1/swap") response = httpx.Response(429, request=request) - error = httpx.HTTPStatusError( - "rate limited", request=request, response=response - ) + error = httpx.HTTPStatusError("rate limited", request=request, response=response) assert classify_redemption_error(error) == ( "mint_unreachable", diff --git a/ui/components/detailed-wallet-balance.tsx b/ui/components/detailed-wallet-balance.tsx index dae15906..3400cf9b 100644 --- a/ui/components/detailed-wallet-balance.tsx +++ b/ui/components/detailed-wallet-balance.tsx @@ -105,6 +105,23 @@ export function DetailedWalletBalance({ const formatMintLabel = (detail: BalanceDetail) => `${detail.mint_url.replace('https://', '').replace('http://', '')} • ${detail.unit.toUpperCase()}`; + const formatBalanceError = (detail: BalanceDetail) => { + const labels: Record = { + rate_limited: 'rate limited', + unreachable: 'unreachable', + cooldown: 'cooling down', + mint_error: 'mint error', + }; + const label = + (detail.error_code ? labels[detail.error_code] : undefined) ?? + detail.error ?? + 'error'; + const retryAfter = detail.retry_after_seconds; + return retryAfter && retryAfter > 0 + ? `${label} (retry in ${Math.ceil(retryAfter)}s)` + : label; + }; + return ( <> @@ -262,9 +279,12 @@ export function DetailedWalletBalance({ {formatMintLabel(detail)} - + {detail.error - ? 'error' + ? formatBalanceError(detail) : formatAmount(walletMsat)} @@ -306,9 +326,12 @@ export function DetailedWalletBalance({

Wallet

-

+

{detail.error - ? 'error' + ? formatBalanceError(detail) : formatAmount(walletMsat)}

diff --git a/ui/lib/api/services/wallet.ts b/ui/lib/api/services/wallet.ts index d16da3ac..cbe25dfc 100644 --- a/ui/lib/api/services/wallet.ts +++ b/ui/lib/api/services/wallet.ts @@ -36,6 +36,8 @@ export interface BalanceDetail { user_balance: number; owner_balance: number; error?: string; + error_code?: 'rate_limited' | 'unreachable' | 'cooldown' | 'mint_error'; + retry_after_seconds?: number; } export interface WithdrawResponse { From 69f19ff991cbfbf64b04d014b7664e332c0b3b3b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 15 Jul 2026 01:43:00 +0200 Subject: [PATCH 19/46] defensive cooldown --- routstr/wallet.py | 59 +++++++++++++++++++++++++++++++++------ tests/unit/test_wallet.py | 27 ++++++++++++++++++ 2 files changed, 78 insertions(+), 8 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index 8828f04c..cbf0552a 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -58,6 +58,8 @@ class TokenConsumedError(Exception): # httpx base classes cover their subclasses. HTTPStatusError is excluded on # purpose — that means the mint answered, just with an error status. _MINT_TRANSPORT_COOLDOWN_SECONDS = 30.0 +_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS = 60.0 +_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS = 7 * 60 * 60 _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( httpx.NetworkError, @@ -90,6 +92,7 @@ class _MintRateGuard: ) self._cooldown_until = 0.0 self._cooldown_reason: str | None = None + self._consecutive_rate_limits = 0 self._needs_probe = False self._probe_lock = asyncio.Lock() @@ -103,6 +106,25 @@ class _MintRateGuard: self._cooldown_reason = reason self._needs_probe = True + def apply_rate_limit_cooldown(self, retry_after: float | None = None) -> float: + remaining = self.cooldown_remaining() + if remaining > 0 and self._cooldown_reason == "rate_limited": + minimum = min( + _MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, + max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0), + ) + if minimum > remaining: + self.apply_cooldown(minimum, reason="rate_limited") + return minimum + return remaining + + self._consecutive_rate_limits += 1 + base = max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0) + multiplier = 2 ** min(self._consecutive_rate_limits - 1, 10) + delay = min(_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, base * multiplier) + self.apply_cooldown(delay, reason="rate_limited") + return delay + def cooldown_remaining(self) -> float: return max(0.0, self._cooldown_until - time.monotonic()) @@ -132,9 +154,16 @@ class _MintRateGuard: try: result = await factory() except Exception as error: - # Keep queued callers behind the probe while _mint_operation applies - # the precise Retry-After/backoff from this failure. - self.apply_cooldown(1.0) + # Keep queued callers behind the probe. Handle rate limits here so + # the next exponential step is recorded before another waiter can + # acquire the probe lock. + if _is_mint_rate_limited(error): + retry_after = None + if isinstance(error, httpx.HTTPStatusError): + retry_after = _parse_retry_after(error.response.headers) + self.apply_rate_limit_cooldown(retry_after) + else: + self.apply_cooldown(1.0) logger.warning( "Mint cooldown probe failed", extra={ @@ -142,6 +171,8 @@ class _MintRateGuard: "mint_url": self._mint_url, "error": str(error), "error_type": type(error).__name__, + "cooldown_seconds": round(self.cooldown_remaining(), 2), + "consecutive_rate_limits": self._consecutive_rate_limits, }, ) raise @@ -149,6 +180,7 @@ class _MintRateGuard: self._needs_probe = False self._cooldown_until = 0.0 self._cooldown_reason = None + self._consecutive_rate_limits = 0 logger.warning( "Mint cooldown probe succeeded; restoring normal concurrency", extra={ @@ -260,8 +292,9 @@ async def _mint_operation( retry_after = _parse_retry_after(exc.response.headers) if retry_after is not None: backoff = max(retry_after, backoff) + cooldown = backoff if guard is not None: - guard.apply_cooldown(backoff, reason="rate_limited") + cooldown = guard.apply_rate_limit_cooldown(backoff) # When the caller has a fallback strategy (trusted-mint # list), re-raise immediately so the caller can try the next @@ -272,7 +305,10 @@ async def _mint_operation( extra={ "op_name": op_name, "mint_url": mint_url, - "cooldown_seconds": round(backoff, 2), + "cooldown_seconds": round(cooldown, 2), + "consecutive_rate_limits": guard._consecutive_rate_limits + if guard is not None + else attempt + 1, }, ) raise @@ -285,11 +321,14 @@ async def _mint_operation( "op_name": op_name, "mint_url": mint_url, "attempt": attempt + 1, - "cooldown_seconds": round(backoff, 2), + "cooldown_seconds": round(cooldown, 2), + "consecutive_rate_limits": guard._consecutive_rate_limits + if guard is not None + else attempt + 1, }, ) if guard is None: - await asyncio.sleep(backoff) + await asyncio.sleep(cooldown) raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") @@ -1648,7 +1687,11 @@ async def fetch_all_balances( if connection_failure else "mint_error" ) - if connection_failure or rate_limited: + if rate_limited: + _MintRateGuard.get(mint_url).apply_rate_limit_cooldown( + _BALANCE_FETCH_RETRY_SECONDS + ) + elif connection_failure: _MintRateGuard.get(mint_url).apply_cooldown( _BALANCE_FETCH_RETRY_SECONDS, reason=error_code ) diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 39d87899..46339065 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1509,6 +1509,33 @@ async def test_mint_rate_guard_waits_for_adaptive_cooldown() -> None: operation.assert_awaited_once() +@pytest.mark.asyncio +async def test_mint_rate_guard_exponentially_backs_off_repeated_429s() -> None: + from routstr.wallet import _MintRateGuard + + guard = _MintRateGuard("http://mint:3338", 4) + expected_delays = [60, 120, 240, 480, 960, 1920, 3840, 7680, 15360, 25200] + now = 0.0 + + with patch("routstr.wallet.time.monotonic") as monotonic: + for index, expected in enumerate(expected_delays, start=1): + monotonic.return_value = now + assert guard.apply_rate_limit_cooldown(60) == expected + assert guard._consecutive_rate_limits == index + if index == 1: + # Concurrent responses from the same 429 wave do not escalate + # the retry count before the first cooldown probe. + assert guard.apply_rate_limit_cooldown(60) == expected + assert guard._consecutive_rate_limits == 1 + now += expected + 1 + + monotonic.return_value = now + operation = AsyncMock(return_value="ok") + assert await guard.run(operation) == "ok" + assert guard._consecutive_rate_limits == 0 + assert guard.apply_rate_limit_cooldown(60) == 60 + + @pytest.mark.asyncio async def test_mint_rate_guard_allows_one_probe_after_cooldown() -> None: from routstr.wallet import _MintRateGuard From 1957e716a377edca77cc053721b7e026c51d5b97 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 15 Jul 2026 01:45:19 +0200 Subject: [PATCH 20/46] clean up keyset unit recog. --- routstr/wallet.py | 12 +++++++----- tests/unit/test_fetch_all_balances.py | 8 ++++---- 2 files changed, 11 insertions(+), 9 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index cbf0552a..5928b492 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -1576,11 +1576,13 @@ async def _get_supported_mint_units(mint_url: str) -> list[str]: mint_url=mint_url, retry_on_rate_limit=False, ) - units = list( - dict.fromkeys( - keyset.unit.name for keyset in keysets if keyset.active and keyset.unit.name - ) - ) + units: list[str] = [] + for keyset in keysets: + if not keyset.active or keyset.unit is None: + continue + unit = keyset.unit if isinstance(keyset.unit, str) else keyset.unit.name + if unit and unit not in units: + units.append(unit) if not units: units = [settings.primary_mint_unit] elif settings.primary_mint_unit in units: diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index ce640e9d..f571aab9 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -121,10 +121,10 @@ async def test_supported_mint_units_come_from_active_keysets() -> None: from routstr.core.settings import settings from routstr.wallet import _get_supported_mint_units - sat = MagicMock(active=True) - sat.unit.name = "sat" - msat = MagicMock(active=False) - msat.unit.name = "msat" + # Cashu versions/mints may deserialize keyset units as either strings or + # Unit enum-like objects. Both representations must be accepted. + sat = MagicMock(active=True, unit="sat") + msat = MagicMock(active=False, unit="msat") usd = MagicMock(active=True) usd.unit.name = "usd" wallet = MagicMock() From 3e906605a0d5faeb2b78f400c330f352fd77da2e Mon Sep 17 00:00:00 2001 From: thefux Date: Sat, 18 Jul 2026 14:34:42 +0000 Subject: [PATCH 21/46] fix: strict rate-limit detection, probe non-escalation, distinct error codes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes three issues that caused per-mint rate-limit state to never recover: 1. _is_mint_rate_limited: remove substring matching on 'rate limit' / 'too many requests' in exception messages. Only HTTP 429 (httpx.HTTPStatusError) is now classified as a rate limit, preventing false positives (e.g. a 503 with 'database rate exceeded' in its body). 2. _run_probe: use apply_cooldown() instead of apply_rate_limit_cooldown() when a probe fails due to a rate limit. The probe is a recovery check, not a new request, so it should not escalate the exponential backoff counter (_consecutive_rate_limits). This prevents the cooldown from ratcheting 60s → 120s → 240s → ... → 7h on repeated probe failures. 3. classify_redemption_error: split the combined _is_mint_rate_limited || is_mint_connection_error check into two separate classifications: - mint_rate_limited / cashu_mint_rate_limited (503, retryable) - mint_unreachable / cashu_mint_unreachable (503, retryable) Callers (routstrd) can now distinguish temporary rate limits from permanent connection failures when deciding fallback strategy. Tests: 20 new tests covering strict 429 detection, classification priority, probe non-escalation, and cooldown reset behaviour. --- routstr/wallet.py | 36 +++++-- tests/unit/test_wallet.py | 193 +++++++++++++++++++++++++++++++++++++- 2 files changed, 217 insertions(+), 12 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index 5928b492..68328f5c 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -154,14 +154,20 @@ class _MintRateGuard: try: result = await factory() except Exception as error: - # Keep queued callers behind the probe. Handle rate limits here so - # the next exponential step is recorded before another waiter can - # acquire the probe lock. + # Keep queued callers behind the probe. On a rate-limit, + # re-apply the *same* cooldown the caller already set rather + # than calling apply_rate_limit_cooldown() — the probe is a + # recovery check, not a new request that should escalate the + # exponential backoff counter. if _is_mint_rate_limited(error): retry_after = None if isinstance(error, httpx.HTTPStatusError): retry_after = _parse_retry_after(error.response.headers) - self.apply_rate_limit_cooldown(retry_after) + delay = max( + _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, + retry_after or 0.0, + ) + self.apply_cooldown(delay, reason="rate_limited") else: self.apply_cooldown(1.0) logger.warning( @@ -217,7 +223,15 @@ def _mint_cooldown_reason(mint_url: str) -> str | None: def _is_mint_rate_limited(error: BaseException) -> bool: - """True if the mint returned a 429 or rate-limit indication.""" + """True if the mint returned an HTTP 429 (Too Many Requests). + + Only matches ``httpx.HTTPStatusError`` with status code 429 — never + classifies based on the exception's message text. Substring matching + on ``"rate limit"`` / ``"too many requests"`` was removed because it + catches unrelated errors (e.g. a 503 whose body happens to mention + "database rate exceeded"), which triggers unnecessary exponential + backoff and can block state recovery indefinitely. + """ current: BaseException | None = error seen: set[int] = set() while current is not None and id(current) not in seen: @@ -225,9 +239,6 @@ def _is_mint_rate_limited(error: BaseException) -> bool: if isinstance(current, httpx.HTTPStatusError): if current.response.status_code == 429: return True - lowered = str(current).lower() - if "rate limit" in lowered or "too many requests" in lowered: - return True current = current.__cause__ or current.__context__ return False @@ -423,7 +434,14 @@ def classify_redemption_error( "The mint that issued this Cashu token is unreachable; the token cannot be redeemed at another mint", "cashu_source_mint_unreachable", ) - if _is_mint_rate_limited(error) or is_mint_connection_error(error): + if _is_mint_rate_limited(error): + return ( + "mint_rate_limited", + 503, + "Cashu mint rate-limited; retry after cooldown", + "cashu_mint_rate_limited", + ) + if is_mint_connection_error(error): return ( "mint_unreachable", 503, diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 46339065..effc03ab 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -12,6 +12,7 @@ from routstr.core.db import ApiKey from routstr.wallet import ( MintConnectionError, TokenConsumedError, + _is_mint_rate_limited, classify_redemption_error, credit_balance, get_balance, @@ -1237,10 +1238,10 @@ def test_rate_limited_mint_is_classified_as_unreachable() -> None: error = httpx.HTTPStatusError("rate limited", request=request, response=response) assert classify_redemption_error(error) == ( - "mint_unreachable", + "mint_rate_limited", 503, - "Cashu mint is unreachable", - "cashu_mint_unreachable", + "Cashu mint rate-limited; retry after cooldown", + "cashu_mint_rate_limited", ) @@ -2126,3 +2127,189 @@ async def test_lightning_fallback_on_429_no_in_place_retry() -> None: assert primary_call_count == 1 assert mock_secondary_wallet.request_mint.await_count == 2 mock_sleep.assert_not_called() + + +# --------------------------------------------------------------------------- +# _is_mint_rate_limited — strict HTTP 429 only (no substring matching) +# --------------------------------------------------------------------------- + +import time as _time_module + + +def _http_429_error(message: str = "") -> httpx.HTTPStatusError: + """Create an HTTP 429 error with optional message in the response body.""" + body = json.dumps({"error": message}) if message else "{}" + return httpx.HTTPStatusError( + message or "Too Many Requests", + request=httpx.Request("POST", "http://m"), + response=httpx.Response(429, content=body.encode()), + ) + + +def _http_500_error(message: str = "") -> httpx.HTTPStatusError: + """Create an HTTP 500 error with optional message in the response body.""" + body = json.dumps({"error": message}) if message else "{}" + return httpx.HTTPStatusError( + message or "Internal Server Error", + request=httpx.Request("POST", "http://m"), + response=httpx.Response(500, content=body.encode()), + ) + + +@pytest.mark.parametrize( + "error,expected", + [ + # True: HTTP 429 is always a rate limit, regardless of message. + (_http_429_error(""), True), + (_http_429_error("Too Many Requests"), True), + (_http_429_error("completely unrelated message"), True), + # False: HTTP 500 is NOT a rate limit, even if the message says "rate limit". + (_http_500_error(""), False), + (_http_500_error("rate limit exceeded"), False), + (_http_500_error("too many requests"), False), + # False: non-HTTP errors with "rate limit" in message. + (ValueError("rate limit exceeded"), False), + (ValueError("too many requests try again"), False), + (RuntimeError("internal rate limit hit"), False), + # False: generic transport errors. + (httpx.ConnectError("connection refused"), False), + (httpx.ReadTimeout("timed out"), False), + (MintConnectionError("mint down"), False), + # Wrapped: HTTP 429 in the cause chain IS detected. + (_chain(ValueError("wrapped"), _http_429_error()), True), + # Wrapped: HTTP 500 with "rate limit" text in cause is NOT detected. + ( + _chain(ValueError("wrapped"), _http_500_error("rate limit exceeded")), + False, + ), + ], +) +def test_is_mint_rate_limited_strictness( + error: BaseException, expected: bool +) -> None: + assert _is_mint_rate_limited(error) is expected + + +def test_is_mint_rate_limited_survives_cycle() -> None: + """A pathological cause/context cycle must not hang the classifier.""" + a = ValueError("a") + b = _http_429_error() + a.__cause__ = b + b.__context__ = a + assert _is_mint_rate_limited(a) is True + + +# --------------------------------------------------------------------------- +# classify_redemption_error — mint_rate_limited vs mint_unreachable +# --------------------------------------------------------------------------- + + +def test_classify_rate_limit_returns_mint_rate_limited() -> None: + """HTTP 429 from a mint is classified as mint_rate_limited, not + mint_unreachable, so callers can distinguish temporary back-off from + permanent mint outages.""" + classified = classify_redemption_error(_http_429_error("Too Many Requests")) + assert classified is not None + type_, status, _msg, code = classified + assert type_ == "mint_rate_limited" + assert status == 503 + assert code == "cashu_mint_rate_limited" + + +def test_classify_rate_limit_takes_priority_over_connection_error() -> None: + """When a 429 is wrapped in a chain that also contains a transport error, + mint_rate_limited wins because it is checked first.""" + inner = _http_429_error() + outer = MintConnectionError("outer") + outer.__cause__ = inner + + classified = classify_redemption_error(outer) + assert classified is not None + type_, status, _msg, code = classified + assert type_ == "mint_rate_limited" + assert code == "cashu_mint_rate_limited" + + +def test_classify_connection_error_still_returns_mint_unreachable() -> None: + """Transport failures without a 429 in the chain are still + classified as mint_unreachable.""" + classified = classify_redemption_error( + httpx.ConnectError("connection refused") + ) + assert classified is not None + type_, status, _msg, code = classified + assert type_ == "mint_unreachable" + assert status == 503 + assert code == "cashu_mint_unreachable" + + +def test_classify_500_with_rate_limit_text_is_not_mint_rate_limited() -> None: + """An HTTP 500 whose body happens to mention 'rate limit' is NOT + classified as mint_rate_limited — it falls through to the generic + error handler.""" + classified = classify_redemption_error( + _http_500_error("database rate limit exceeded") + ) + # Should NOT be mint_rate_limited or mint_unreachable. + if classified is not None: + type_, _status, _msg, code = classified + assert type_ != "mint_rate_limited" + assert code != "cashu_mint_rate_limited" + + +# --------------------------------------------------------------------------- +# _MintRateGuard — probe does NOT escalate cooldown counter +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_probe_does_not_escalate_consecutive_rate_limits() -> None: + """When a probe fails with a rate limit, _consecutive_rate_limits should + NOT increment — the probe is a recovery check, not a new request.""" + from routstr.wallet import _MintRateGuard, _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS + + guard = _MintRateGuard("http://mint", max_concurrency=0) + + # Simulate initial rate limit: apply_rate_limit_cooldown increments counter + guard.apply_rate_limit_cooldown() + assert guard._consecutive_rate_limits == 1 + cooldown_before = guard._cooldown_until + assert cooldown_before > 0 + + # Simulate probe failure: _run_probe uses apply_cooldown, NOT + # apply_rate_limit_cooldown, so the counter stays at 1. + guard.apply_cooldown( + _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, reason="rate_limited" + ) + assert guard._consecutive_rate_limits == 1 # unchanged! + assert guard._needs_probe is True + + +@pytest.mark.asyncio +async def test_probe_recovery_resets_consecutive_rate_limits() -> None: + """A successful probe resets _consecutive_rate_limits to 0.""" + from routstr.wallet import _MintRateGuard + + guard = _MintRateGuard("http://mint", max_concurrency=0) + + # First rate limit: increments to 1, sets 60s cooldown. + guard.apply_rate_limit_cooldown() + assert guard._consecutive_rate_limits == 1 + + # Manually expire the cooldown so the next call creates a fresh one. + guard._cooldown_until = 0.0 + guard._cooldown_reason = None + + # Second rate limit (after cooldown expired): increments to 2. + guard.apply_rate_limit_cooldown() + assert guard._consecutive_rate_limits == 2 + + # Simulate a successful probe by resetting (as _run_probe does) + guard._needs_probe = False + guard._cooldown_until = 0.0 + guard._cooldown_reason = None + guard._consecutive_rate_limits = 0 + + assert guard._consecutive_rate_limits == 0 + assert guard._needs_probe is False + assert guard.cooldown_remaining() == 0.0 From 586af15a1bf4fef976ea38c9ebd95b120cb52f0d Mon Sep 17 00:00:00 2001 From: thefux Date: Sat, 18 Jul 2026 14:41:32 +0000 Subject: [PATCH 22/46] chore: fix ruff lint errors (E402, F401, I001) --- tests/unit/test_wallet.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index effc03ab..108db581 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -2133,8 +2133,6 @@ async def test_lightning_fallback_on_429_no_in_place_retry() -> None: # _is_mint_rate_limited — strict HTTP 429 only (no substring matching) # --------------------------------------------------------------------------- -import time as _time_module - def _http_429_error(message: str = "") -> httpx.HTTPStatusError: """Create an HTTP 429 error with optional message in the response body.""" @@ -2266,7 +2264,7 @@ def test_classify_500_with_rate_limit_text_is_not_mint_rate_limited() -> None: async def test_probe_does_not_escalate_consecutive_rate_limits() -> None: """When a probe fails with a rate limit, _consecutive_rate_limits should NOT increment — the probe is a recovery check, not a new request.""" - from routstr.wallet import _MintRateGuard, _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS + from routstr.wallet import _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, _MintRateGuard guard = _MintRateGuard("http://mint", max_concurrency=0) From a9a638161427371345bac3cf4c13d7aa3af9b62f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 22 Jul 2026 22:27:50 +0200 Subject: [PATCH 23/46] fix: recreate mint URL migration from latest head --- ...ab843b49_add_mint_url_to_lightning_invoices.py} | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) rename migrations/versions/{21c84cd5ad83_add_mint_url_to_lightning_invoices.py => 11eaab843b49_add_mint_url_to_lightning_invoices.py} (52%) diff --git a/migrations/versions/21c84cd5ad83_add_mint_url_to_lightning_invoices.py b/migrations/versions/11eaab843b49_add_mint_url_to_lightning_invoices.py similarity index 52% rename from migrations/versions/21c84cd5ad83_add_mint_url_to_lightning_invoices.py rename to migrations/versions/11eaab843b49_add_mint_url_to_lightning_invoices.py index b69c16a8..40142d8f 100644 --- a/migrations/versions/21c84cd5ad83_add_mint_url_to_lightning_invoices.py +++ b/migrations/versions/11eaab843b49_add_mint_url_to_lightning_invoices.py @@ -1,22 +1,24 @@ """add mint url to lightning invoices -Revision ID: 21c84cd5ad83 -Revises: c6d7e8f9a0b1 -Create Date: 2026-07-12 15:04:01.675455 +Revision ID: 11eaab843b49 +Revises: d7e8f9a0b1c2 +Create Date: 2026-07-22 22:25:45.278261 """ import sqlalchemy as sa from alembic import op # revision identifiers, used by Alembic. -revision = "21c84cd5ad83" -down_revision = "c6d7e8f9a0b1" +revision = "11eaab843b49" +down_revision = "d7e8f9a0b1c2" branch_labels = None depends_on = None def upgrade() -> None: - op.add_column("lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True)) + op.add_column( + "lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True) + ) def downgrade() -> None: From 92246b78d0bdce1ecce75e09ad3173e1c8bd3916 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 23 Jul 2026 01:08:34 +0200 Subject: [PATCH 24/46] add import --- routstr/proxy.py | 1 + 1 file changed, 1 insertion(+) diff --git a/routstr/proxy.py b/routstr/proxy.py index 7156671f..27e3372a 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -19,6 +19,7 @@ from .core.db import ( ) from .core.exceptions import UpstreamError from .core.not_found import build_not_found_response +from .core.settings import settings from .payment.helpers import ( apply_mint_fee_allowance, calculate_discounted_max_cost, From ab80657507e921766dcba568be64b5d7d699fff2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 20:56:34 +0200 Subject: [PATCH 25/46] recreate model paths migration --- .../bda277ee4683_add_model_paths_table.py | 55 ++++++++++++++++ .../d7e8f9a0b1c2_add_model_paths_table.py | 66 ------------------- 2 files changed, 55 insertions(+), 66 deletions(-) create mode 100644 migrations/versions/bda277ee4683_add_model_paths_table.py delete mode 100644 migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py diff --git a/migrations/versions/bda277ee4683_add_model_paths_table.py b/migrations/versions/bda277ee4683_add_model_paths_table.py new file mode 100644 index 00000000..ab5ce131 --- /dev/null +++ b/migrations/versions/bda277ee4683_add_model_paths_table.py @@ -0,0 +1,55 @@ +"""add model paths table + +Revision ID: bda277ee4683 +Revises: fc4fa29630d2 +Create Date: 2026-07-24 20:54:56.822687 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "bda277ee4683" +down_revision = "fc4fa29630d2" +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.ForeignKeyConstraint( + ["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE" + ), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + op.create_index( + op.f("ix_model_paths_model_id"), "model_paths", ["model_id"], unique=False + ) + op.create_index( + op.f("ix_model_paths_upstream_provider_id"), + "model_paths", + ["upstream_provider_id"], + unique=False, + ) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths") + op.drop_index(op.f("ix_model_paths_model_id"), table_name="model_paths") + op.drop_table("model_paths") + # ### end Alembic commands ### diff --git a/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py b/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py deleted file mode 100644 index e5ea85f4..00000000 --- a/migrations/versions/d7e8f9a0b1c2_add_model_paths_table.py +++ /dev/null @@ -1,66 +0,0 @@ -"""add model_paths table - -Revision ID: d7e8f9a0b1c2 -Revises: c6d7e8f9a0b1 -Create Date: 2026-07-05 00:00:00.000000 -""" - -from __future__ import annotations - -import sqlalchemy as sa -from alembic import op - -revision = "d7e8f9a0b1c2" -down_revision = "c6d7e8f9a0b1" -branch_labels = None -depends_on = None - - -def upgrade() -> None: - conn = op.get_bind() - inspector = sa.inspect(conn) - if "model_paths" in inspector.get_table_names(): - return - - op.create_table( - "model_paths", - sa.Column("id", sa.Integer(), primary_key=True, autoincrement=True), - sa.Column("model_id", sa.String(), nullable=False), - sa.Column("path", sa.String(), nullable=False), - sa.Column("upstream_provider_id", sa.Integer(), nullable=False), - sa.ForeignKeyConstraint( - ["upstream_provider_id"], - ["upstream_providers.id"], - ondelete="CASCADE", - ), - sa.UniqueConstraint( - "model_id", - "path", - "upstream_provider_id", - name="uq_model_paths_model_path_provider", - ), - ) - op.create_index( - "ix_model_paths_model_id", - "model_paths", - ["model_id"], - ) - op.create_index( - "ix_model_paths_upstream_provider_id", - "model_paths", - ["upstream_provider_id"], - ) - - -def downgrade() -> None: - conn = op.get_bind() - inspector = sa.inspect(conn) - if "model_paths" not in inspector.get_table_names(): - return - - existing_indexes = {idx["name"] for idx in inspector.get_indexes("model_paths")} - if "ix_model_paths_upstream_provider_id" in existing_indexes: - op.drop_index("ix_model_paths_upstream_provider_id", table_name="model_paths") - if "ix_model_paths_model_id" in existing_indexes: - op.drop_index("ix_model_paths_model_id", table_name="model_paths") - op.drop_table("model_paths") From 4c292580e8d25d8239d828c8af7cc3448d1336d1 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 21:15:55 +0200 Subject: [PATCH 26/46] rebase model paths migration onto latest head --- ..._table.py => 4e0c3d195a49_add_model_paths_table.py} | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) rename migrations/versions/{bda277ee4683_add_model_paths_table.py => 4e0c3d195a49_add_model_paths_table.py} (91%) diff --git a/migrations/versions/bda277ee4683_add_model_paths_table.py b/migrations/versions/4e0c3d195a49_add_model_paths_table.py similarity index 91% rename from migrations/versions/bda277ee4683_add_model_paths_table.py rename to migrations/versions/4e0c3d195a49_add_model_paths_table.py index ab5ce131..61710f64 100644 --- a/migrations/versions/bda277ee4683_add_model_paths_table.py +++ b/migrations/versions/4e0c3d195a49_add_model_paths_table.py @@ -1,8 +1,8 @@ """add model paths table -Revision ID: bda277ee4683 -Revises: fc4fa29630d2 -Create Date: 2026-07-24 20:54:56.822687 +Revision ID: 4e0c3d195a49 +Revises: 7f2843d3f4e4 +Create Date: 2026-07-24 21:14:39.062179 """ import sqlalchemy as sa @@ -10,8 +10,8 @@ import sqlmodel from alembic import op # revision identifiers, used by Alembic. -revision = "bda277ee4683" -down_revision = "fc4fa29630d2" +revision = "4e0c3d195a49" +down_revision = "7f2843d3f4e4" branch_labels = None depends_on = None From 81c0ff57e92a0115a8045bdc8785fe61cc7d65a2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 24 Jul 2026 22:24:24 +0200 Subject: [PATCH 27/46] fix: expose cost breakdown in paid responses --- routstr/payment/cost_calculation.py | 30 +++- routstr/upstream/base.py | 140 +++++++++++++++++-- tests/unit/test_cost_calculation_caching.py | 4 + tests/unit/test_messages_litellm_dispatch.py | 3 +- tests/unit/test_x_cashu_cost_sats.py | 12 +- 5 files changed, 175 insertions(+), 14 deletions(-) diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 37ac15d3..d89f70ee 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -481,12 +481,38 @@ def _calculate_from_usd_cost( ) output_msats = cost_in_msats - input_msats + # Estimate cache read/creation msats proportionally within the input cost. + # These are informational subcomponents: input_msats remains inclusive of + # cache cost so input_msats + output_msats == total_msats, matching the + # token-priced path and the public CostData contract. + cache_read_msats = 0 + cache_creation_msats = 0 + if cache_read_tokens > 0 or cache_creation_tokens > 0: + cache_tokens = cache_read_tokens + cache_creation_tokens + regular_input_tokens = input_tokens + total_input_tokens = regular_input_tokens + cache_tokens + if total_input_tokens > 0: + # Approximate by token count because the USD path only exposes an + # aggregate input cost, not separately priced cache buckets. + cache_read_msats = ( + int(input_msats * cache_read_tokens / total_input_tokens) + if cache_read_tokens > 0 + else 0 + ) + cache_creation_msats = ( + int(input_msats * cache_creation_tokens / total_input_tokens) + if cache_creation_tokens > 0 + else 0 + ) + logger.info( "Using cost from usage data/details", extra={ "usd_cost": usd_cost, "cost_in_sats": cost_in_sats, "cost_in_msats": cost_in_msats, + "cache_read_msats": cache_read_msats, + "cache_creation_msats": cache_creation_msats, "model": response_data.get("model", "unknown"), }, ) @@ -501,8 +527,8 @@ def _calculate_from_usd_cost( output_tokens=output_tokens, cache_read_input_tokens=cache_read_tokens, cache_creation_input_tokens=cache_creation_tokens, - cache_read_msats=0, - cache_creation_msats=0, + cache_read_msats=cache_read_msats, + cache_creation_msats=cache_creation_msats, ) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index a8dba7e3..343230cf 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -69,6 +69,56 @@ if typing.TYPE_CHECKING: logger = get_logger(__name__) +def _inject_cost_response_headers( + headers: dict[str, str], cost_data: CostData | MaxCostData +) -> None: + """Inject per-request cost breakdown into response headers. + + The SDK's ``extractUsageFromResponseHeaders`` reads these to populate + ``inputMsats``, ``outputMsats``, ``totalMsats`` and ``satsCost`` in the + usage tracking entry — without them, x-cashu requests show 0.0 for all + sat cost fields. + """ + headers["X-Routstr-Cost-Msats"] = str(cost_data.total_msats) + headers["X-Routstr-Input-Cost-Msats"] = str(cost_data.input_msats) + headers["X-Routstr-Output-Cost-Msats"] = str(cost_data.output_msats) + if cost_data.total_usd: + headers["X-Routstr-Cost-Usd"] = str(cost_data.total_usd) + + +def _inject_cost_into_usage( + response_json: dict, cost_data: CostData | MaxCostData +) -> None: + """Inject cost breakdown into the response body's ``usage.cost`` object. + + The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be + an object with ``total_msats``/``input_msats``/``output_msats`` (not a + plain USD number). When the upstream returns ``cost`` as a number, the + SDK cannot extract the msats breakdown from the body alone. + """ + usage = response_json.get("usage") + if not isinstance(usage, dict): + return + # Direct assignment (not setdefault) so routstr's authoritative cost + # data always overwrites any upstream-provided cost values. Using + # setdefault would silently keep stale upstream values and drop our + # calculated msats breakdown. + cost_obj: dict[str, int | float] = { + "base_msats": cost_data.base_msats, + "input_msats": cost_data.input_msats, + "output_msats": cost_data.output_msats, + "total_msats": cost_data.total_msats, + "cache_read_input_tokens": cost_data.cache_read_input_tokens, + "cache_creation_input_tokens": cost_data.cache_creation_input_tokens, + "cache_read_msats": cost_data.cache_read_msats, + "cache_creation_msats": cost_data.cache_creation_msats, + } + if cost_data.total_usd: + cost_obj["total_usd"] = cost_data.total_usd + usage["cost"] = cost_obj + usage["cost_sats"] = cost_data.total_msats // 1000 + + def _is_json_content_type(content_type: str | None) -> bool: """Return True when the upstream response should be parsed as JSON.""" if not content_type: @@ -284,9 +334,27 @@ class BaseUpstreamProvider: sats_cost = total_msats // 1000 + # Build the cost object that the SDK's extractUsageFromResponseBody + # and extractUsageFromSSEJson expect: an object with total_msats, + # input_msats, output_msats, cache_read_msats, cache_creation_msats, + # etc. Setting usage.cost to a plain float (total_usd) means the SDK + # cannot extract the msats breakdown — cache_read_msats and + # cache_creation_msats in particular are lost. + cost_obj = { + "base_msats": cost_dict.get("base_msats", 0), + "input_msats": cost_dict.get("input_msats", 0), + "output_msats": cost_dict.get("output_msats", 0), + "total_msats": total_msats, + "total_usd": total_usd, + "cache_read_input_tokens": cost_dict.get("cache_read_input_tokens", 0), + "cache_creation_input_tokens": cost_dict.get("cache_creation_input_tokens", 0), + "cache_read_msats": cost_dict.get("cache_read_msats", 0), + "cache_creation_msats": cost_dict.get("cache_creation_msats", 0), + } + # Inject into top-level usage block (OpenAI/Anthropic style) if "usage" in response_json: - response_json["usage"]["cost"] = total_usd + response_json["usage"]["cost"] = cost_obj response_json["usage"]["cost_sats"] = sats_cost response_json["usage"]["remaining_balance_msats"] = key.balance self._fold_cache_into_input_tokens(response_json["usage"]) @@ -2153,6 +2221,21 @@ class BaseUpstreamProvider: if k.lower() in allowed_headers } + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for balance-mode requests. + if isinstance(cost_data, dict): + _cost_data_obj = CostData( + base_msats=cost_data.get("base_msats", 0), + input_msats=cost_data.get("input_msats", 0), + output_msats=cost_data.get("output_msats", 0), + total_msats=cost_data.get("total_msats", 0), + total_usd=cost_data.get("total_usd", 0.0), + ) + else: + _cost_data_obj = cost_data + _inject_cost_response_headers(response_headers, _cost_data_obj) + return Response( content=json.dumps(response_json).encode(), status_code=response.status_code, @@ -2242,9 +2325,24 @@ class BaseUpstreamProvider: ) self.inject_cost_metadata(response_json, cost_data, key) + # Inject cost breakdown headers for balance-mode requests. + if isinstance(cost_data, dict): + _cost_data_obj = CostData( + base_msats=cost_data.get("base_msats", 0), + input_msats=cost_data.get("input_msats", 0), + output_msats=cost_data.get("output_msats", 0), + total_msats=cost_data.get("total_msats", 0), + total_usd=cost_data.get("total_usd", 0.0), + ) + else: + _cost_data_obj = cost_data + response_headers: dict[str, str] = {} + _inject_cost_response_headers(response_headers, _cost_data_obj) + return Response( content=json.dumps(response_json).encode(), status_code=200, + headers=response_headers, media_type="application/json", ) @@ -2295,11 +2393,12 @@ class BaseUpstreamProvider: and "usage" in response_json and isinstance(response_json["usage"], dict) ): - response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + _inject_cost_into_usage(response_json, cost_data) self._fold_cache_into_input_tokens(response_json["usage"]) response_headers: dict[str, str] = {} if cost_data: + _inject_cost_response_headers(response_headers, cost_data) refund_amount = messages_dispatch.compute_refund( amount, unit, cost_data.total_msats ) @@ -3700,6 +3799,11 @@ class BaseUpstreamProvider: "model": model, }, ) + + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for x-cashu requests. + _inject_cost_response_headers(response_headers, cost_data) except Exception as e: logger.error( "Error calculating cost for streaming response", @@ -3722,8 +3826,12 @@ class BaseUpstreamProvider: if "provider" not in data_json: self._apply_provider_field(data_json) changed = True - if cost_data and "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + if ( + cost_data + and "usage" in data_json + and data_json["usage"] + ): + _inject_cost_into_usage(data_json, cost_data) changed = True if changed: lines[i] = "data: " + json.dumps(data_json) @@ -3777,7 +3885,10 @@ class BaseUpstreamProvider: ) if cost_data and "usage" in response_json: - response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + # Inject cost breakdown into both the response body (so the + # SDK's body extractor picks up the msats breakdown) and the + # response headers (so the SDK's header extractor works too). + _inject_cost_into_usage(response_json, cost_data) if not cost_data: logger.error( @@ -3808,6 +3919,8 @@ class BaseUpstreamProvider: if "content-encoding" in response_headers: del response_headers["content-encoding"] + _inject_cost_response_headers(response_headers, cost_data) + if unit == "msat": refund_amount = amount - cost_data.total_msats elif unit == "sat": @@ -4681,6 +4794,11 @@ class BaseUpstreamProvider: "model": model, }, ) + + # Inject cost breakdown headers so the SDK's + # extractUsageFromResponseHeaders can populate + # inputMsats/outputMsats/totalMsats for x-cashu requests. + _inject_cost_response_headers(response_headers, cost_data) except Exception as e: logger.error( "Error calculating cost for streaming Responses API response", @@ -4703,8 +4821,12 @@ class BaseUpstreamProvider: if "provider" not in data_json: self._apply_provider_field(data_json) changed = True - if cost_data and "usage" in data_json and data_json["usage"]: - data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + if ( + cost_data + and "usage" in data_json + and data_json["usage"] + ): + _inject_cost_into_usage(data_json, cost_data) changed = True if changed: lines[i] = "data: " + json.dumps(data_json) @@ -4747,7 +4869,7 @@ class BaseUpstreamProvider: ) if cost_data and "usage" in response_json: - response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 + _inject_cost_into_usage(response_json, cost_data) if not cost_data: logger.error( @@ -4778,6 +4900,8 @@ class BaseUpstreamProvider: if "content-encoding" in response_headers: del response_headers["content-encoding"] + _inject_cost_response_headers(response_headers, cost_data) + if unit == "msat": refund_amount = amount - cost_data.total_msats elif unit == "sat": diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index ba31366a..0722e5bf 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -529,6 +529,8 @@ async def test_openrouter_upstream_inference_cost_components_are_used() -> None: assert isinstance(result, CostData) assert result.input_msats == 994 assert result.output_msats == 3477 + assert result.cache_read_msats == 758 + assert result.cache_creation_msats == 0 assert result.input_msats + result.output_msats == result.total_msats == 4471 @@ -574,6 +576,8 @@ async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None: # Token normalisation (OpenAI dialect: cached included in prompt_tokens) assert result.input_tokens == 5070 # 164371 - 159301 assert result.cache_read_input_tokens == 159301 + assert result.cache_read_msats == 897966 + assert result.cache_creation_msats == 0 assert result.output_tokens == 99 diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index c41568ad..79d16030 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -451,7 +451,8 @@ async def test_non_streaming_dispatches_via_litellm_and_returns_anthropic_respon assert payload["model"] == "openai/gpt-4o-mini" # mapped back to requested assert payload["usage"]["input_tokens"] == 5 assert payload["usage"]["output_tokens"] == 3 - assert payload["usage"]["cost"] == 0.0001 + assert payload["usage"]["cost"]["total_msats"] == 1234 + assert payload["usage"]["cost"]["total_usd"] == 0.0001 assert payload["usage"]["cost_sats"] == 1 diff --git a/tests/unit/test_x_cashu_cost_sats.py b/tests/unit/test_x_cashu_cost_sats.py index 0dc509cf..901cc2f6 100644 --- a/tests/unit/test_x_cashu_cost_sats.py +++ b/tests/unit/test_x_cashu_cost_sats.py @@ -67,8 +67,13 @@ async def test_non_streaming_includes_cost_sats() -> None: ) body = json.loads(response.body) - assert "cost_sats" in body["usage"] assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000 + assert body["usage"]["cost"]["total_msats"] == 5000 + assert body["usage"]["cost"]["input_msats"] == 3000 + assert body["usage"]["cost"]["output_msats"] == 2000 + assert response.headers["x-routstr-cost-msats"] == "5000" + assert response.headers["x-routstr-input-cost-msats"] == "3000" + assert response.headers["x-routstr-output-cost-msats"] == "2000" @pytest.mark.asyncio @@ -96,7 +101,7 @@ async def test_non_streaming_cost_sats_value_rounds_down() -> None: @pytest.mark.asyncio -async def test_non_streaming_preserves_existing_usage_fields() -> None: +async def test_non_streaming_preserves_tokens_and_replaces_upstream_cost() -> None: provider = _make_provider() cost_data = _make_cost_data(total_msats=3000) @@ -127,7 +132,8 @@ async def test_non_streaming_preserves_existing_usage_fields() -> None: assert usage["prompt_tokens"] == 100 assert usage["completion_tokens"] == 50 assert usage["total_tokens"] == 150 - assert usage["cost"] == 0.00015 + assert usage["cost"]["total_msats"] == 3000 + assert usage["cost"]["total_usd"] == 0.00025 assert usage["cost_sats"] == 3 From 4c6bc49e072327bf9001d64f4f216476e18ece6a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 25 Jul 2026 00:36:32 +0200 Subject: [PATCH 28/46] unify payment path --- routstr/payment/cost_calculation.py | 81 ++++++--- routstr/upstream/base.py | 181 +++++++++---------- tests/unit/test_cost_calculation_caching.py | 132 +++++++++++++- tests/unit/test_cost_response_metadata.py | 113 ++++++++++++ tests/unit/test_messages_litellm_dispatch.py | 6 + 5 files changed, 395 insertions(+), 118 deletions(-) create mode 100644 tests/unit/test_cost_response_metadata.py diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index d89f70ee..0fb305cb 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -183,6 +183,28 @@ async def calculate_cost( cost_details.get("output_cost") or cost_details.get("upstream_inference_completions_cost") ) + cache_pricing_rates: tuple[float, float, float, float] | None = None + if cache_read_tokens > 0 or cache_creation_tokens > 0: + try: + cache_pricing_rates = _get_pricing_rates( + response_data, model_obj, provider_fee + ) + except ValueError: + logger.warning( + "Cache pricing unavailable for USD cost breakdown; " + "leaving cache cost components unknown", + extra={"model": response_data.get("model", "unknown")}, + ) + if cache_pricing_rates is None and settings.fixed_pricing: + fixed_input_rate = ( + float(settings.fixed_per_1k_input_tokens) * 1000.0 + ) + cache_pricing_rates = ( + fixed_input_rate, + float(settings.fixed_per_1k_output_tokens) * 1000.0, + fixed_input_rate, + fixed_input_rate, + ) return _calculate_from_usd_cost( usd_cost, input_usd, @@ -193,6 +215,7 @@ async def calculate_cost( output_tokens, response_data, provider_fee, + cache_pricing_rates, ) except Exception as e: logger.warning( @@ -451,6 +474,7 @@ def _calculate_from_usd_cost( output_tokens: int, response_data: dict, provider_fee: float | None, + pricing_rates: tuple[float, float, float, float] | None = None, ) -> CostData: """Calculate cost from USD figures, deriving input/output split from tokens.""" if provider_fee is None: @@ -460,15 +484,20 @@ def _calculate_from_usd_cost( output_usd = output_usd * provider_fee sats_per_usd = 1.0 / sats_usd_price() cost_in_sats = usd_cost * sats_per_usd - cost_in_msats = math.ceil(cost_in_sats * 1000) + raw_cost_msats = cost_in_sats * 1000 + cost_in_msats = math.ceil(raw_cost_msats) + raw_input_msats = 0.0 if input_usd > 0 or output_usd > 0: # The total is the authoritative billed amount. Allocating that integer # total proportionally avoids losing sub-millisatoshi remainders when # input and output components are each truncated independently. component_usd = input_usd + output_usd - input_msats = math.floor(cost_in_msats * input_usd / component_usd) - output_msats = cost_in_msats - input_msats + # Match the token-priced path: truncate the visible output component + # and assign the authoritative total's rounding remainder to input. + output_msats = math.floor(cost_in_msats * output_usd / component_usd) + input_msats = cost_in_msats - output_msats + raw_input_msats = raw_cost_msats * input_usd / component_usd else: effective_input_tokens = ( input_tokens + cache_read_tokens + cache_creation_tokens @@ -480,29 +509,37 @@ def _calculate_from_usd_cost( else 0 ) output_msats = cost_in_msats - input_msats + raw_input_msats = ( + raw_cost_msats * effective_input_tokens / total_tokens + if total_tokens > 0 + else 0.0 + ) - # Estimate cache read/creation msats proportionally within the input cost. - # These are informational subcomponents: input_msats remains inclusive of - # cache cost so input_msats + output_msats == total_msats, matching the - # token-priced path and the public CostData contract. + # Preserve the same cache-rate ratios as the token-priced path while the + # upstream USD total remains authoritative. Cache values are informational + # subcomponents of the inclusive input cost. cache_read_msats = 0 cache_creation_msats = 0 - if cache_read_tokens > 0 or cache_creation_tokens > 0: - cache_tokens = cache_read_tokens + cache_creation_tokens - regular_input_tokens = input_tokens - total_input_tokens = regular_input_tokens + cache_tokens - if total_input_tokens > 0: - # Approximate by token count because the USD path only exposes an - # aggregate input cost, not separately priced cache buckets. - cache_read_msats = ( - int(input_msats * cache_read_tokens / total_input_tokens) - if cache_read_tokens > 0 - else 0 + if pricing_rates is not None: + input_rate, _, cache_read_rate, cache_creation_rate = pricing_rates + regular_weight = input_tokens * input_rate + cache_read_weight = cache_read_tokens * cache_read_rate + cache_creation_weight = cache_creation_tokens * cache_creation_rate + total_input_weight = ( + regular_weight + cache_read_weight + cache_creation_weight + ) + if total_input_weight > 0: + cache_read_msats = int( + round( + raw_input_msats * cache_read_weight / total_input_weight, + 3, + ) ) - cache_creation_msats = ( - int(input_msats * cache_creation_tokens / total_input_tokens) - if cache_creation_tokens > 0 - else 0 + cache_creation_msats = int( + round( + raw_input_msats * cache_creation_weight / total_input_weight, + 3, + ) ) logger.info( diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 343230cf..3353df8a 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -69,8 +69,21 @@ if typing.TYPE_CHECKING: logger = get_logger(__name__) +CostMetadata = CostData | MaxCostData | dict[str, Any] + + +def _cost_field( + cost_data: CostMetadata, field: str, default: int | float = 0 +) -> int | float: + if isinstance(cost_data, dict): + value = cost_data.get(field, default) + else: + value = getattr(cost_data, field, default) + return value if isinstance(value, (int, float)) else default + + def _inject_cost_response_headers( - headers: dict[str, str], cost_data: CostData | MaxCostData + headers: dict[str, str], cost_data: CostMetadata ) -> None: """Inject per-request cost breakdown into response headers. @@ -79,16 +92,21 @@ def _inject_cost_response_headers( usage tracking entry — without them, x-cashu requests show 0.0 for all sat cost fields. """ - headers["X-Routstr-Cost-Msats"] = str(cost_data.total_msats) - headers["X-Routstr-Input-Cost-Msats"] = str(cost_data.input_msats) - headers["X-Routstr-Output-Cost-Msats"] = str(cost_data.output_msats) - if cost_data.total_usd: - headers["X-Routstr-Cost-Usd"] = str(cost_data.total_usd) + headers["X-Routstr-Cost-Msats"] = str( + int(_cost_field(cost_data, "total_msats")) + ) + headers["X-Routstr-Input-Cost-Msats"] = str( + int(_cost_field(cost_data, "input_msats")) + ) + headers["X-Routstr-Output-Cost-Msats"] = str( + int(_cost_field(cost_data, "output_msats")) + ) + total_usd = float(_cost_field(cost_data, "total_usd", 0.0)) + if total_usd: + headers["X-Routstr-Cost-Usd"] = str(total_usd) -def _inject_cost_into_usage( - response_json: dict, cost_data: CostData | MaxCostData -) -> None: +def _inject_cost_into_usage(response_json: dict, cost_data: CostMetadata) -> None: """Inject cost breakdown into the response body's ``usage.cost`` object. The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be @@ -104,19 +122,26 @@ def _inject_cost_into_usage( # setdefault would silently keep stale upstream values and drop our # calculated msats breakdown. cost_obj: dict[str, int | float] = { - "base_msats": cost_data.base_msats, - "input_msats": cost_data.input_msats, - "output_msats": cost_data.output_msats, - "total_msats": cost_data.total_msats, - "cache_read_input_tokens": cost_data.cache_read_input_tokens, - "cache_creation_input_tokens": cost_data.cache_creation_input_tokens, - "cache_read_msats": cost_data.cache_read_msats, - "cache_creation_msats": cost_data.cache_creation_msats, + "base_msats": int(_cost_field(cost_data, "base_msats")), + "input_msats": int(_cost_field(cost_data, "input_msats")), + "output_msats": int(_cost_field(cost_data, "output_msats")), + "total_msats": int(_cost_field(cost_data, "total_msats")), + "cache_read_input_tokens": int( + _cost_field(cost_data, "cache_read_input_tokens") + ), + "cache_creation_input_tokens": int( + _cost_field(cost_data, "cache_creation_input_tokens") + ), + "cache_read_msats": int(_cost_field(cost_data, "cache_read_msats")), + "cache_creation_msats": int( + _cost_field(cost_data, "cache_creation_msats") + ), } - if cost_data.total_usd: - cost_obj["total_usd"] = cost_data.total_usd + total_usd = float(_cost_field(cost_data, "total_usd", 0.0)) + if total_usd: + cost_obj["total_usd"] = total_usd usage["cost"] = cost_obj - usage["cost_sats"] = cost_data.total_msats // 1000 + usage["cost_sats"] = int(_cost_field(cost_data, "total_msats")) // 1000 def _is_json_content_type(content_type: str | None) -> bool: @@ -325,48 +350,24 @@ class BaseUpstreamProvider: self._apply_provider_field(response_json) if isinstance(cost_data, dict): total_msats = cost_data.get("total_msats", 0) - total_usd = cost_data.get("total_usd", 0.0) cost_dict = cost_data else: total_msats = cost_data.total_msats - total_usd = cost_data.total_usd cost_dict = cost_data.dict() sats_cost = total_msats // 1000 - # Build the cost object that the SDK's extractUsageFromResponseBody - # and extractUsageFromSSEJson expect: an object with total_msats, - # input_msats, output_msats, cache_read_msats, cache_creation_msats, - # etc. Setting usage.cost to a plain float (total_usd) means the SDK - # cannot extract the msats breakdown — cache_read_msats and - # cache_creation_msats in particular are lost. - cost_obj = { - "base_msats": cost_dict.get("base_msats", 0), - "input_msats": cost_dict.get("input_msats", 0), - "output_msats": cost_dict.get("output_msats", 0), - "total_msats": total_msats, - "total_usd": total_usd, - "cache_read_input_tokens": cost_dict.get("cache_read_input_tokens", 0), - "cache_creation_input_tokens": cost_dict.get("cache_creation_input_tokens", 0), - "cache_read_msats": cost_dict.get("cache_read_msats", 0), - "cache_creation_msats": cost_dict.get("cache_creation_msats", 0), - } - - # Inject into top-level usage block (OpenAI/Anthropic style) - if "usage" in response_json: - response_json["usage"]["cost"] = cost_obj - response_json["usage"]["cost_sats"] = sats_cost + # Inject the shared SDK cost contract into every usage shape. + if isinstance(response_json.get("usage"), dict): + _inject_cost_into_usage(response_json, cost_data) response_json["usage"]["remaining_balance_msats"] = key.balance self._fold_cache_into_input_tokens(response_json["usage"]) - # Inject into Anthropic nested usage block if present - if ( - "message" in response_json - and isinstance(response_json["message"], dict) - and "usage" in response_json["message"] - ): - response_json["message"]["usage"]["sats_cost"] = sats_cost - self._fold_cache_into_input_tokens(response_json["message"]["usage"]) + message = response_json.get("message") + if isinstance(message, dict) and isinstance(message.get("usage"), dict): + _inject_cost_into_usage(message, cost_data) + message["usage"]["remaining_balance_msats"] = key.balance + self._fold_cache_into_input_tokens(message["usage"]) # Unified Routstr metadata response_json["metadata"] = response_json.get("metadata", {}) @@ -1297,12 +1298,9 @@ class BaseUpstreamProvider: await session.refresh(key) remaining_balance_msats = key.balance - # Merge cost into usage for OpenCode + # Merge the shared cost contract into usage for SDKs and OpenCode. if "usage" in response_json: - response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0) - response_json["usage"]["cost_sats"] = ( - cost_data.get("total_msats", 0) // 1000 - ) + _inject_cost_into_usage(response_json, cost_data) response_json["usage"]["remaining_balance_msats"] = ( remaining_balance_msats ) @@ -1349,6 +1347,7 @@ class BaseUpstreamProvider: for k, v in response.headers.items() if k.lower() in allowed_headers } + _inject_cost_response_headers(response_headers, cost_data) if requested_model: response_json["model"] = requested_model @@ -1734,12 +1733,9 @@ class BaseUpstreamProvider: await session.refresh(key) remaining_balance_msats = key.balance - # Merge cost into usage for OpenCode + # Merge the shared cost contract into usage for SDKs and OpenCode. if "usage" in response_json: - response_json["usage"]["cost"] = cost_data.get("total_usd", 0.0) - response_json["usage"]["cost_sats"] = ( - cost_data.get("total_msats", 0) // 1000 - ) + _inject_cost_into_usage(response_json, cost_data) response_json["usage"]["remaining_balance_msats"] = ( remaining_balance_msats ) @@ -1786,6 +1782,7 @@ class BaseUpstreamProvider: for k, v in response.headers.items() if k.lower() in allowed_headers } + _inject_cost_response_headers(response_headers, cost_data) if requested_model: response_json["model"] = requested_model @@ -2221,20 +2218,8 @@ class BaseUpstreamProvider: if k.lower() in allowed_headers } - # Inject cost breakdown headers so the SDK's - # extractUsageFromResponseHeaders can populate - # inputMsats/outputMsats/totalMsats for balance-mode requests. - if isinstance(cost_data, dict): - _cost_data_obj = CostData( - base_msats=cost_data.get("base_msats", 0), - input_msats=cost_data.get("input_msats", 0), - output_msats=cost_data.get("output_msats", 0), - total_msats=cost_data.get("total_msats", 0), - total_usd=cost_data.get("total_usd", 0.0), - ) - else: - _cost_data_obj = cost_data - _inject_cost_response_headers(response_headers, _cost_data_obj) + # Inject the same cost headers used by every paid response path. + _inject_cost_response_headers(response_headers, cost_data) return Response( content=json.dumps(response_json).encode(), @@ -2325,19 +2310,9 @@ class BaseUpstreamProvider: ) self.inject_cost_metadata(response_json, cost_data, key) - # Inject cost breakdown headers for balance-mode requests. - if isinstance(cost_data, dict): - _cost_data_obj = CostData( - base_msats=cost_data.get("base_msats", 0), - input_msats=cost_data.get("input_msats", 0), - output_msats=cost_data.get("output_msats", 0), - total_msats=cost_data.get("total_msats", 0), - total_usd=cost_data.get("total_usd", 0.0), - ) - else: - _cost_data_obj = cost_data + # Inject the same cost headers used by every paid response path. response_headers: dict[str, str] = {} - _inject_cost_response_headers(response_headers, _cost_data_obj) + _inject_cost_response_headers(response_headers, cost_data) return Response( content=json.dumps(response_json).encode(), @@ -2644,7 +2619,7 @@ class BaseUpstreamProvider: the cost of a wire-format change for clients that read ``X-Cashu`` from headers today. """ - buffered: list[bytes] = [] + buffered: list[messages_dispatch.AnnotatedEvent] = [] last_model_seen: str | None = None input_tokens = 0 output_tokens = 0 @@ -2672,7 +2647,7 @@ class BaseUpstreamProvider: total_cost = max(total_cost, annotated.total_cost) input_cost = max(input_cost, annotated.input_cost) output_cost = max(output_cost, annotated.output_cost) - buffered.append(annotated.sse_bytes) + buffered.append(annotated) response_headers: dict[str, str] = { "Cache-Control": "no-cache", @@ -2699,6 +2674,7 @@ class BaseUpstreamProvider: }, ) + cost_data: CostData | MaxCostData | None = None if ( input_tokens > 0 or output_tokens > 0 @@ -2754,9 +2730,30 @@ class BaseUpstreamProvider: }, ) + if cost_data: + _inject_cost_response_headers(response_headers, cost_data) + for index, annotated in enumerate(buffered): + event = annotated.event + changed = False + message = event.get("message") + if isinstance(message, dict) and isinstance(message.get("usage"), dict): + _inject_cost_into_usage(message, cost_data) + changed = True + if isinstance(event.get("usage"), dict): + _inject_cost_into_usage(event, cost_data) + changed = True + if changed: + event_type = str(event.get("type") or "") + prefix = f"event: {event_type}\n" if event_type else "" + buffered[index] = annotated._replace( + sse_bytes=( + f"{prefix}data: {json.dumps(event)}\n\n".encode() + ) + ) + async def replay() -> AsyncGenerator[bytes, None]: - for chunk in buffered: - yield chunk + for annotated in buffered: + yield annotated.sse_bytes return StreamingResponse( replay(), diff --git a/tests/unit/test_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index 0722e5bf..65ad5091 100644 --- a/tests/unit/test_cost_calculation_caching.py +++ b/tests/unit/test_cost_calculation_caching.py @@ -15,6 +15,7 @@ os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to") from routstr.core.settings import settings from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost +from routstr.payment.models import Architecture, Model, Pricing @pytest.fixture(autouse=True) @@ -527,13 +528,136 @@ async def test_openrouter_upstream_inference_cost_components_are_used() -> None: result = await calculate_cost(response, max_cost=100000) assert isinstance(result, CostData) - assert result.input_msats == 994 - assert result.output_msats == 3477 + assert result.input_msats == 995 + assert result.output_msats == 3476 assert result.cache_read_msats == 758 assert result.cache_creation_msats == 0 assert result.input_msats + result.output_msats == result.total_msats == 4471 +@pytest.mark.asyncio +async def test_usd_cache_breakdown_matches_token_priced_path( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Authoritative USD totals must retain model-specific cache-rate ratios.""" + monkeypatch.setattr(settings, "fixed_pricing", False) + model = Model( + id="cache-priced-model", + name="cache-priced-model", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test", + instruct_type=None, + ), + pricing=Pricing(prompt=0.01, completion=0.02), + sats_pricing=Pricing( + prompt=0.01, + completion=0.02, + input_cache_read=0.001, + input_cache_write=0.01, + ), + per_request_limits=None, + top_provider=None, + ) + usage = { + "prompt_tokens": 1000, + "completion_tokens": 100, + "prompt_tokens_details": {"cached_tokens": 900}, + } + + token_result = await calculate_cost( + {"model": model.id, "usage": usage}, + max_cost=100_000, + model_obj=model, + ) + usd_result = await calculate_cost( + { + "model": model.id, + "usage": { + **usage, + "cost": 0.000195, + "cost_details": { + "input_cost": 0.000095, + "output_cost": 0.0001, + }, + }, + }, + max_cost=100_000, + model_obj=model, + provider_fee=1.0, + ) + + assert isinstance(token_result, CostData) + assert isinstance(usd_result, CostData) + assert usd_result.total_msats == token_result.total_msats == 3900 + assert usd_result.input_msats + usd_result.output_msats == usd_result.total_msats + assert usd_result.cache_read_msats == token_result.cache_read_msats == 900 + + +@pytest.mark.asyncio +async def test_usd_cache_breakdown_does_not_absorb_total_rounding_remainder( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Sub-msat cache components truncate like the token-priced path.""" + monkeypatch.setattr(settings, "fixed_pricing", False) + model = Model( + id="sub-msat-cache-model", + name="sub-msat-cache-model", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="test", + instruct_type=None, + ), + pricing=Pricing(prompt=0.001, completion=0.001), + sats_pricing=Pricing( + prompt=0.001, + completion=0.001, + input_cache_write=0.0006, + ), + per_request_limits=None, + top_provider=None, + ) + usage = { + "input_tokens": 0, + "output_tokens": 0, + "cache_creation_input_tokens": 1, + } + + token_result = await calculate_cost( + {"model": model.id, "usage": usage}, + max_cost=100_000, + model_obj=model, + ) + usd_result = await calculate_cost( + { + "model": model.id, + "usage": { + **usage, + "cost": 0.00000003, + "cost_details": {"input_cost": 0.00000003}, + }, + }, + max_cost=100_000, + model_obj=model, + provider_fee=1.0, + ) + + assert isinstance(token_result, CostData) + assert isinstance(usd_result, CostData) + assert usd_result.total_msats == token_result.total_msats == 1 + assert usd_result.cache_creation_msats == token_result.cache_creation_msats == 0 + + # ============================================================================ # PPQ.AI BYOK: upstream_inference_cost + BYOK fee billing # @@ -570,8 +694,8 @@ async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None: # msats), not the fee alone (~0.0023 USD → ~45k msats). ~20× correction. assert result.total_msats == 940274 assert result.input_msats + result.output_msats == result.total_msats - assert result.input_msats == 926546 - assert result.output_msats == 13728 + assert result.input_msats == 926547 + assert result.output_msats == 13727 assert result.total_usd == pytest.approx(0.047013667305) # Token normalisation (OpenAI dialect: cached included in prompt_tokens) assert result.input_tokens == 5070 # 164371 - 159301 diff --git a/tests/unit/test_cost_response_metadata.py b/tests/unit/test_cost_response_metadata.py new file mode 100644 index 00000000..beaf4005 --- /dev/null +++ b/tests/unit/test_cost_response_metadata.py @@ -0,0 +1,113 @@ +"""Response-contract tests for Routstr cost metadata across paid paths.""" + +import json +import os +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.db import ApiKey # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 + +COST_DATA = { + "base_msats": 0, + "input_msats": 1_200, + "output_msats": 300, + "total_msats": 1_500, + "total_usd": 0.0001, + "input_tokens": 10, + "output_tokens": 3, + "cache_read_input_tokens": 8, + "cache_creation_input_tokens": 2, + "cache_read_msats": 80, + "cache_creation_msats": 40, +} + + +def _provider() -> BaseUpstreamProvider: + return BaseUpstreamProvider(base_url="http://test", api_key="upstream-key") + + +def _key() -> ApiKey: + return ApiKey(hashed_key="abcdef0123" * 4, balance=1_000_000) + + +def _session() -> Any: + session = MagicMock() + session.refresh = AsyncMock() + return session + + +def _upstream_response(payload: dict) -> httpx.Response: + return httpx.Response( + 200, + json=payload, + request=httpx.Request("POST", "http://test"), + ) + + +def _assert_cost_contract(response: Any) -> None: + body = json.loads(response.body) + assert body["usage"]["cost"] == { + "base_msats": 0, + "input_msats": 1_200, + "output_msats": 300, + "total_msats": 1_500, + "total_usd": 0.0001, + "cache_read_input_tokens": 8, + "cache_creation_input_tokens": 2, + "cache_read_msats": 80, + "cache_creation_msats": 40, + } + assert response.headers["X-Routstr-Cost-Msats"] == "1500" + assert response.headers["X-Routstr-Input-Cost-Msats"] == "1200" + assert response.headers["X-Routstr-Output-Cost-Msats"] == "300" + + +@pytest.mark.asyncio +async def test_balance_chat_completion_uses_shared_cost_contract() -> None: + provider = _provider() + with patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(return_value=dict(COST_DATA)), + ): + response = await provider.handle_non_streaming_chat_completion( + _upstream_response( + { + "model": "test-model", + "usage": {"prompt_tokens": 10, "completion_tokens": 3}, + } + ), + _key(), + _session(), + deducted_max_cost=10_000, + ) + + _assert_cost_contract(response) + + +@pytest.mark.asyncio +async def test_balance_responses_completion_uses_shared_cost_contract() -> None: + provider = _provider() + with patch( + "routstr.upstream.base.adjust_payment_for_tokens", + new=AsyncMock(return_value=dict(COST_DATA)), + ): + response = await provider.handle_non_streaming_responses_completion( + _upstream_response( + { + "model": "test-model", + "usage": {"input_tokens": 10, "output_tokens": 3}, + } + ), + _key(), + _session(), + deducted_max_cost=10_000, + ) + + _assert_cost_contract(response) diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index 79d16030..cb3fe9d7 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -855,6 +855,9 @@ async def test_x_cashu_streaming_replays_events_and_sets_refund_header() -> None assert isinstance(result, StreamingResponse) assert result.headers.get("X-Cashu") == "cashuSTREAM" + assert result.headers.get("X-Routstr-Cost-Msats") == "1500000" + assert result.headers.get("X-Routstr-Input-Cost-Msats") == "1000000" + assert result.headers.get("X-Routstr-Output-Cost-Msats") == "500000" # 1_500_000 msats → 1500 sats. Refund = 5000 - 1500 = 3500. mock_refund.assert_awaited_once() refund_call = mock_refund.await_args @@ -873,6 +876,9 @@ async def test_x_cashu_streaming_replays_events_and_sets_refund_header() -> None assert "event: message_start" in joined assert "event: message_delta" in joined assert "event: message_stop" in joined + assert '"total_msats": 1500000' in joined + assert '"input_msats": 1000000' in joined + assert '"output_msats": 500000' in joined # --------------------------------------------------------------------------- From 0a00527626202b5960c1152b66e8caa903b3f580 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:23:11 +0200 Subject: [PATCH 29/46] chore: apply ruff format repo-wide CI only runs ruff check, so format drift accumulated. Committed separately so the reformat noise stays out of functional commits. --- routstr/algorithm.py | 20 ++++++++-------- routstr/balance.py | 29 +++++++++++++++++++---- routstr/core/admin.py | 26 ++++++++++----------- routstr/core/log_manager.py | 21 ++++++++--------- routstr/core/usage_analytics_store.py | 29 +++++++++++------------ routstr/nostr/analytics.py | 8 +++++-- routstr/payment/cost_calculation.py | 19 ++++----------- routstr/payment/usage.py | 4 +--- routstr/upstream/azure.py | 4 +--- routstr/upstream/ehbp.py | 32 +++++++++++++++----------- routstr/upstream/gemini.py | 4 +--- routstr/upstream/gemini_messages.py | 4 +--- routstr/upstream/groq.py | 4 +++- routstr/upstream/litellm_routing.py | 4 +--- routstr/upstream/messages_dispatch.py | 12 +++------- routstr/upstream/ollama.py | 8 +++---- routstr/upstream/rate_limit.py | 4 +++- routstr/upstream/request_correction.py | 4 +--- routstr/upstream/routstr.py | 3 +-- routstr/upstream/xai.py | 4 +++- routstr/wallet.py | 13 ++++++++--- 21 files changed, 132 insertions(+), 124 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index fbc5388e..ef4b8574 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -232,7 +232,10 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: + if ( + model_to_use.forwarded_model_id + and model_to_use.forwarded_model_id not in aliases + ): aliases.append(model_to_use.forwarded_model_id) # Try to set each alias @@ -322,7 +325,10 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: + if ( + model_to_use.forwarded_model_id + and model_to_use.forwarded_model_id not in aliases + ): aliases.append(model_to_use.forwarded_model_id) for alias in aliases: @@ -342,16 +348,10 @@ def create_model_mappings( forwarded_model_ids, the one whose forwarded_model_id equals the requested alias wins. """ - if ( - model.forwarded_model_id - and model.forwarded_model_id.lower() == alias - ): + if model.forwarded_model_id and model.forwarded_model_id.lower() == alias: return 5 - if ( - model.id - and model.id.lower() == alias - ): + if model.id and model.id.lower() == alias: return 4 model_base = get_base_model_id(model.id) diff --git a/routstr/balance.py b/routstr/balance.py index 91b19ce5..03dc4d33 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -260,7 +260,11 @@ async def _lookup_key_no_create( async def _restore_balance( - session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str + session: AsyncSession, + hashed_key: str, + balance: int, + reserved_balance: int, + mint_url: str, ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -275,7 +279,11 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, + extra={ + "hashed_key": hashed_key, + "restored_balance": balance, + "mint_url": mint_url, + }, ) @@ -460,11 +468,23 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") + await _restore_balance( + session, + key.hashed_key, + pre_debit_balance, + pre_debit_reserved, + key.refund_mint_url or "", + ) error_msg = str(e) logger.error( "refund_wallet_endpoint: mint/send failed", @@ -685,7 +705,6 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} - @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 66a1d288..1521510e 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -68,7 +68,9 @@ async def require_admin_api(request: Request) -> None: async with create_session() as session: result = await session.exec(select(CliToken).where(CliToken.token == token)) cli_token = result.first() - if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts): + if cli_token and ( + cli_token.expires_at is None or cli_token.expires_at > now_ts + ): cli_token.last_used_at = now_ts session.add(cli_token) await session.commit() @@ -255,16 +257,12 @@ async def update_password(request: Request, password_update: PasswordUpdate) -> secret = await get_secret(session) if not secret.admin_password_hash: - raise HTTPException( - status_code=500, detail="Admin password not configured" - ) + raise HTTPException(status_code=500, detail="Admin password not configured") if not vault.verify_password( password_update.current_password, secret.admin_password_hash ): - raise HTTPException( - status_code=401, detail="Current password is incorrect" - ) + raise HTTPException(status_code=401, detail="Current password is incorrect") # Validate new password new_password = password_update.new_password.strip() @@ -980,9 +978,7 @@ async def update_upstream_provider_by_slug( lookup = _validate_slug(payload.slug) async with create_session() as session: result = await session.exec( - select(UpstreamProviderRow).where( - UpstreamProviderRow.slug == lookup - ) + select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup) ) provider = result.first() if not provider: @@ -1669,7 +1665,11 @@ async def get_transactions_api( ) total = count_result.one() - stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit) + stmt = ( + base.order_by(col(CashuTransaction.created_at).desc()) + .offset(offset) + .limit(limit) + ) results = await session.exec(stmt) transactions = results.all() @@ -1679,9 +1679,7 @@ async def get_transactions_api( } -@admin_router.get( - "/api/lightning-invoices", dependencies=[Depends(require_admin_api)] -) +@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)]) async def get_lightning_invoices_api( status: str | None = None, purpose: str | None = None, diff --git a/routstr/core/log_manager.py b/routstr/core/log_manager.py index 0444dcbf..b111f68a 100644 --- a/routstr/core/log_manager.py +++ b/routstr/core/log_manager.py @@ -408,7 +408,9 @@ class LogManager: def get_error_details(self, hours: int = 24, limit: int = 100) -> dict: def compute() -> dict: try: - return self._usage_store.get_error_details(hours_back=hours, limit=limit) + return self._usage_store.get_error_details( + hours_back=hours, limit=limit + ) except Exception as e: logger.error( f"Usage analytics index failed, falling back to log scan: {e}" @@ -628,8 +630,7 @@ class LogManager: stats["total_tokens"] += input_tokens + output_tokens failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: stats["total_requests"] += 1 @@ -787,7 +788,9 @@ class LogManager: if bucket_key: model_mix_buckets[bucket_key][model] += 1 if revenue_msats > 0: - model_mix_revenue_buckets[bucket_key][model] += revenue_msats + model_mix_revenue_buckets[bucket_key][model] += ( + revenue_msats + ) model_mix_revenue_totals[model] += revenue_msats if input_tokens > 0 or output_tokens > 0: token_total = input_tokens + output_tokens @@ -801,8 +804,7 @@ class LogManager: bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: summary_stats["total_requests"] += 1 @@ -872,9 +874,7 @@ class LogManager: models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True) latest_errors = [ item - for _, item in sorted( - latest_errors_heap, key=lambda x: x[0], reverse=True - ) + for _, item in sorted(latest_errors_heap, key=lambda x: x[0], reverse=True) ] top_model_limit = max(1, min(model_limit, 20)) top_models_requests = [ @@ -1051,8 +1051,7 @@ class LogManager: bucket["warnings"] += 1 failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: bucket["total_requests"] += 1 diff --git a/routstr/core/usage_analytics_store.py b/routstr/core/usage_analytics_store.py index 7ba90e24..36fa4bcb 100644 --- a/routstr/core/usage_analytics_store.py +++ b/routstr/core/usage_analytics_store.py @@ -314,9 +314,7 @@ class UsageAnalyticsStore: if column in existing_columns: return - conn.execute( - f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}" - ) + conn.execute(f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}") logger.info(f"Migrated analytics schema: added {table}.{column}") def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None: @@ -364,7 +362,11 @@ class UsageAnalyticsStore: self._drop_index_tables_locked(conn) self._initialize_schema_locked(conn) - files = log_files if log_files is not None else sorted(self.logs_dir.glob("app_*.log")) + files = ( + log_files + if log_files is not None + else sorted(self.logs_dir.glob("app_*.log")) + ) for log_file in files: try: self._process_log_file_locked(conn, log_file, force_full_read=True) @@ -568,8 +570,7 @@ class UsageAnalyticsStore: model_bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message - or "revert payment" in message + "upstream request failed" in message or "revert payment" in message ) if failed: bucket["total_requests"] += 1 @@ -592,9 +593,9 @@ class UsageAnalyticsStore: if isinstance(max_cost, (int, float)) and max_cost > 0: max_cost_float = float(max_cost) bucket["refunds_msats"] += max_cost_float - model_updates[(minute_key, model)][ - "refunds_msats" - ] += max_cost_float + model_updates[(minute_key, model)]["refunds_msats"] += ( + max_cost_float + ) return ( end_offset, @@ -1032,7 +1033,9 @@ class UsageAnalyticsStore: """, (cutoff_timestamp,), ).fetchone() - total_error_count = int(total_error_count_row[0]) if total_error_count_row else 0 + total_error_count = ( + int(total_error_count_row[0]) if total_error_count_row else 0 + ) return { "errors": [ @@ -1204,11 +1207,7 @@ class UsageAnalyticsStore: total_successful = int(row["total_successful"]) total_revenue_msats = float(row["total_revenue_msats"]) total_tokens = int(row["total_tokens"]) - if ( - total_successful <= 0 - and total_revenue_msats <= 0 - and total_tokens <= 0 - ): + if total_successful <= 0 and total_revenue_msats <= 0 and total_tokens <= 0: continue bucket_ts = str(row["bucket_ts"]) diff --git a/routstr/nostr/analytics.py b/routstr/nostr/analytics.py index e568b5e0..8b6da590 100644 --- a/routstr/nostr/analytics.py +++ b/routstr/nostr/analytics.py @@ -215,7 +215,9 @@ def _build_window_payload( summary = dashboard.get("summary", {}) model_usage_mix = dashboard.get("model_usage_mix", {}) - summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {}) + summary_payload = _build_summary_payload( + summary if isinstance(summary, dict) else {} + ) usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {} top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload) @@ -338,7 +340,9 @@ async def publish_usage_analytics() -> None: nsec = (settings.nsec or "").strip() if not nsec: if not warned_missing_nsec: - logger.info("NSEC is not configured; skipping analytics sharing to Nostr") + logger.info( + "NSEC is not configured; skipping analytics sharing to Nostr" + ) warned_missing_nsec = True await asyncio.sleep(DISABLED_POLL_SECONDS) continue diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 37ac15d3..e7cee8ca 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -224,9 +224,7 @@ async def calculate_cost( "Token counts %s in the upstream response but cannot be " "priced; the request will appear in dashboards with the " "raw counts and a fixed max-cost charge.", - "are present" - if (input_tokens > 0 or output_tokens > 0) - else "are zero", + "are present" if (input_tokens > 0 or output_tokens > 0) else "are zero", extra={ "base_cost_msats": max_cost, "model": response_data.get("model", "unknown"), @@ -303,9 +301,7 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: # actually deducts from the balance. For non-BYOK providers (e.g. # OpenRouter) usage.cost already equals upstream_inference_cost, so we # fall through to the normal ``cost`` lookup below. - upstream_cost = _coerce_usd( - cost_details.get("upstream_inference_cost") - ) + upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost")) if upstream_cost > 0 and usage_data.get("is_byok"): byok_fee = _coerce_usd(usage_data.get("cost")) return upstream_cost + byok_fee @@ -336,8 +332,7 @@ def _get_pricing_rates( ``None`` means configured fixed pricing should be used by the caller. """ if settings.fixed_pricing and ( - settings.fixed_per_1k_input_tokens - or settings.fixed_per_1k_output_tokens + settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens ): return None @@ -393,12 +388,8 @@ def _get_pricing_rates( usd_per_sat = sats_usd_price() mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat - cache_read_usd = _coerce_usd( - pricing.get("cache_read_input_token_cost") - ) - cache_write_usd = _coerce_usd( - pricing.get("cache_creation_input_token_cost") - ) + cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost")) + cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost")) mscr_1k = ( cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat if cache_read_usd > 0 diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 02c90055..11d5c01e 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -110,9 +110,7 @@ def normalize_usage(usage_data: object) -> NormalizedUsage | None: if not isinstance(usage_data, dict): return None - output_tokens = _first_token_count( - usage_data, "completion_tokens", "output_tokens" - ) + output_tokens = _first_token_count(usage_data, "completion_tokens", "output_tokens") cache_read, cache_write = _extract_cache_tokens(usage_data) # ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index a693b763..985bcfd2 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -94,9 +94,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider): deployment_id = deployment_id.split("/")[-1] return f"openai/deployments/{deployment_id}/{clean_path}" - def get_request_base_url( - self, path: str, model_obj: "Model | None" = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: """Use endpoint root, stripping accidental /openai/v1 suffix if present.""" base_url = self.base_url.rstrip("/") marker = "/openai/v1" diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 96955492..a3d1505a 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -191,7 +191,9 @@ def _resolve_ehbp_target_url( otherwise the header is ignored so callers cannot redirect other providers or leak upstream API keys. """ - override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER + override_header = ( + profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER + ) if not override_header: return target_url enclave_url = _get_header_case_insensitive(headers, override_header) @@ -295,9 +297,7 @@ def _build_cost_info( return result -def _inject_cost_response_headers( - headers: dict[str, str], cost_info: dict -) -> None: +def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None: """Add per-request cost headers to an EHBP response. Since EHBP response bodies are opaque encrypted blobs, cost cannot be @@ -375,9 +375,7 @@ async def _compute_ehbp_actual_cost( resolved_upstream_model = ( actual_model_obj.forwarded_model_id or actual_model_obj.id ) - resolved_identity = _normalize_upstream_model_id( - resolved_upstream_model - ) + resolved_identity = _normalize_upstream_model_id(resolved_upstream_model) if resolved_identity != expected_identity: logger.info( "EHBP served model differs from requested, using actual " @@ -517,7 +515,9 @@ async def finalize_ehbp_actual_cost_payment( billing_key = await get_billing_key(key, session) key_hash = key.hashed_key billing_key_hash = billing_key.hashed_key - total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model))) + total_cost_msats = max( + 0, int(cost_info.get("total_msats", reserved_cost_for_model)) + ) now = int(time.time()) safe_reserved = case( @@ -560,7 +560,9 @@ async def finalize_ehbp_actual_cost_payment( ) child_result = await session.exec(child_stmt) # type: ignore[call-overload] - if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): + if result.rowcount == 0 or ( + child_result is not None and child_result.rowcount == 0 + ): await session.rollback() logger.error( "Failed to finalize EHBP usage-based payment", @@ -690,7 +692,9 @@ async def finalize_ehbp_max_cost_payment( else: child_result = None - if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): + if result.rowcount == 0 or ( + child_result is not None and child_result.rowcount == 0 + ): await session.rollback() logger.error( "Failed to finalize EHBP max-cost payment", @@ -1034,7 +1038,9 @@ async def forward_ehbp_x_cashu_request( target_url = _resolve_ehbp_target_url( target.url, path, headers, provider_type, profile ) - upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile) + upstream_headers = _prepare_ehbp_upstream_headers( + headers, target.headers, profile + ) request_body = await request.body() # Merge query params into the target URL @@ -1082,9 +1088,7 @@ async def forward_ehbp_x_cashu_request( usage_source = ( "header" if usage_header_name - and any( - k.lower() == usage_header_name.lower() for k, _ in resp.headers - ) + and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers) else ("trailer" if usage_header else "none") ) diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index 54de41a3..d58199c2 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -94,9 +94,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): """ return self.base_url.rstrip("/").removesuffix("/openai") + "/openai" - def get_request_base_url( - self, path: str, model_obj: "Model | None" = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: """Route every proxied request to the OpenAI-compat surface. Required because the stored ``base_url`` typically points at the diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index 11440b91..a7e41c9b 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -371,9 +371,7 @@ async def dispatch_gemini_messages( aggregates). """ if not request_body: - raise UpstreamError( - "Missing request body for /v1/messages", status_code=400 - ) + raise UpstreamError("Missing request body for /v1/messages", status_code=400) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index 17103c35..a0c9475e 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -20,7 +20,9 @@ class GroqUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": + def _build_from_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "GroqUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/upstream/litellm_routing.py b/routstr/upstream/litellm_routing.py index 0b2a92a3..b7790394 100644 --- a/routstr/upstream/litellm_routing.py +++ b/routstr/upstream/litellm_routing.py @@ -91,9 +91,7 @@ OLLAMA_HOST_HINTS: tuple[str, ...] = ( ) -def detect_litellm_prefix( - base_url: str | None, default: str = DEFAULT_PREFIX -) -> str: +def detect_litellm_prefix(base_url: str | None, default: str = DEFAULT_PREFIX) -> str: """Return the litellm provider prefix (`"/"`) for `base_url`. Falls back to `default` when the host doesn't match any known provider. diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 3d689922..efcf591b 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -108,9 +108,7 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: return events, buffer -def events_from_chunk( - chunk: object, sse_buffer: bytes -) -> tuple[list[dict], bytes]: +def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]: """Normalize a stream chunk into one or more event dicts. ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE @@ -201,9 +199,7 @@ async def aggregate_anthropic_events_to_message( raw_json = partial_json.pop(idx, None) if raw_json is not None and idx < len(blocks): try: - blocks[idx]["input"] = ( - json.loads(raw_json) if raw_json else {} - ) + blocks[idx]["input"] = json.loads(raw_json) if raw_json else {} except json.JSONDecodeError: blocks[idx]["input"] = raw_json elif etype == "message_delta": @@ -445,9 +441,7 @@ async def dispatch_anthropic_messages( on bad input or upstream failure. """ if not request_body: - raise UpstreamError( - "Missing request body for /v1/messages", status_code=400 - ) + raise UpstreamError("Missing request body for /v1/messages", status_code=400) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index 9fed0154..c4873ea0 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -66,9 +66,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - def get_request_base_url( - self, path: str, model_obj: Model | None = None - ) -> str: + def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" return f"{self.base_url.rstrip('/')}/v1" @@ -185,7 +183,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): except Exception: self._models_cache = models_with_fees - self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache} + self._models_by_id = { + m.forwarded_model_id or m.id: m for m in self._models_cache + } logger.info( f"Refreshed models cache for {self.base_url}", extra={"model_count": len(models)}, diff --git a/routstr/upstream/rate_limit.py b/routstr/upstream/rate_limit.py index dca1ba5b..ac1eff78 100644 --- a/routstr/upstream/rate_limit.py +++ b/routstr/upstream/rate_limit.py @@ -119,7 +119,9 @@ def classify_rate_limit( retry_match = _RETRY_RE.search(redacted) if retry_match is not None: value = float(retry_match.group(1)) - retry_after = value / 1000.0 if retry_match.group(2).lower() == "ms" else value + retry_after = ( + value / 1000.0 if retry_match.group(2).lower() == "ms" else value + ) limit_name_match = _LIMIT_NAME_RE.search(redacted) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index c2ea5b1d..8e4379a0 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -84,9 +84,7 @@ def extract_error_message(response: Response) -> str: return "" -def strip_unsupported_param( - body: dict, error_message: str -) -> tuple[dict, str] | None: +def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: """Drop a top-level param the upstream named as unsupported/deprecated. Returns ``(new_body, param)`` (a new dict, original untouched) when the diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index 0371946a..de1aa3bd 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -50,8 +50,7 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): def normalize_request_path( self, path: str, model_obj: "Model | None" = None ) -> str: - """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr. - """ + """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.""" return path.lstrip("/") @classmethod diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 58caaba0..12e3dd93 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -21,7 +21,9 @@ class XAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": + def _build_from_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "XAIUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/wallet.py b/routstr/wallet.py index dd92d913..cef6902d 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -243,7 +243,9 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { - f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id) + f"{k.mint_url}/{k.unit.name}": sum( + p.amount for p in wallet.proofs if p.id == k.id + ) for k in wallet.keysets.values() } # Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet @@ -598,11 +600,16 @@ async def swap_to_primary_mint( # advance the counter so the next request derives fresh secrets. logger.warning( "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", - extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, + extra={ + "mint_quote_id": mint_quote.quote, + "minted_amount": minted_amount, + }, ) try: for keyset_id in primary_wallet.keysets: - await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) + await primary_wallet.restore_tokens_for_keyset( + keyset_id, to=1, batch=25 + ) await primary_wallet.load_proofs(reload=True) post_recovery_balance = primary_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance From f96acbb99cc5dbc11e64f9f933d8520031aaf059 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:23:31 +0200 Subject: [PATCH 30/46] fix: address model-paths review findings Provider scoping (items 1/2/6): - Key visibility maps on (model_id.lower(), upstream_provider_id), matching refresh_model_maps, so a disable/override row on one provider never leaks onto another provider's model, and matching is case-insensitive. Data safety (items 3/5): - Degraded OpenRouter fetches (network error, 429, non-200, bad payload) return None (unknown) instead of []; a provider whose path set is unknown keeps its previously persisted rows instead of being wiped. - Endpoint payload parsing moved fully inside try, with a list guard, so endpoints:null or non-list shapes are swallowed as documented. - refresh with an empty live upstream list is a no-op; the unfiltered DELETE in the prune path is gone (prune now keys off enabled DB rows). Hot path (items 4/12/14): - Persist uses chunked bulk INSERTs (one statement per 500 rows) instead of per-row ORM adds; redundant ix_model_paths_model_id index dropped. - Read routes filter in SQL instead of materializing the whole table, and output ordering is deterministic (public id + path), independent of rowid. - Visibility no longer rebuilds fully priced Model objects per override row; it reads id/forwarded_model_id/canonical_slug straight off ModelRow. Path/id contract (items 7/8/9/11): - discovery_path_for_subprovider/discovery_base_paths hooks on BaseUpstreamProvider, overridden by OpenRouterUpstreamProvider, mirror _apply_provider_field so discovery and response stamping cannot drift (openrouter:OpenRouter now correctly maps to unknown). - openrouter_author_slug falls back to a slash-containing forwarded_model_id, so admin-created alias rows are discoverable. - public_model_id splits on the first slash, same as get_base_model_id, so discovery ids can be sent to chat completions verbatim. Lifecycle (items 10/13): - ENABLE_MODEL_PATHS_REFRESH kill switch; interval and flag re-read every loop iteration, and the task idles (not exits) while disabled. - First 429 latches and aborts the remaining fan-out for the cycle; a per-cycle cache dedupes fetches across providers sharing a base URL. - refresh_model_maps prunes paths of disabled/deleted providers so admin mutations take effect immediately; rows carry updated_at and both endpoints expose it. Tests (item 15) rewritten through the public refresh entry point with transport-level httpx.MockTransport fakes, FK enforcement on, and coverage for the periodic loop. Migration re-chained onto 9c4d8e2f1a6b. --- docs/api/endpoints.md | 11 +- docs/provider/configuration.md | 3 +- .../4e0c3d195a49_add_model_paths_table.py | 15 +- routstr/core/db.py | 25 +- routstr/core/main.py | 23 +- routstr/core/settings.py | 18 +- routstr/payment/models.py | 12 +- routstr/proxy.py | 13 + routstr/upstream/base.py | 22 + routstr/upstream/model_paths.py | 469 ++++--- routstr/upstream/openrouter.py | 17 + tests/unit/test_fee_payout_migration.py | 4 +- tests/unit/test_model_paths.py | 1103 ++++++++++++----- 13 files changed, 1183 insertions(+), 552 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index efe57a68..44e6207e 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -374,10 +374,13 @@ GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 } ``` -Model IDs in responses are unqualified display IDs: provider prefixes such as -`z-ai/` or `openai/` are stripped. Path values match the provider string stamped -on chat-completion responses, such as `anthropic`, `generic:my-upstream`, or -`openrouter:Anthropic`. +Model IDs in responses are base model IDs: the leading provider prefix such as +`z-ai/` or `openai/` is stripped (the same rule routing uses, so the ID can be +sent back to `/v1/chat/completions` verbatim). Path values match the provider +string stamped on chat-completion responses, such as `anthropic`, +`generic:Anthropic`, `openrouter:Anthropic`, or `unknown` (native OpenRouter +with no usable sub-provider). Responses also carry an `updated_at` Unix +timestamp of the last successful refresh (`null` when no refresh has run). ## Wallet Management diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 6dd9d338..8bcbca9b 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -142,7 +142,8 @@ Use environment variables for: | `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` | | `CORS_ORIGINS` | Allowed CORS origins | `*` | | `RELAYS` | Nostr relays (comma-separated) | (default set) | -| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to disable | `600` | +| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to pause the refresh (previously discovered paths keep being served) | `600` | +| `ENABLE_MODEL_PATHS_REFRESH` | Kill switch for the background model-path refresh (OpenRouter endpoint fan-out) | `true` | ### Priority diff --git a/migrations/versions/4e0c3d195a49_add_model_paths_table.py b/migrations/versions/4e0c3d195a49_add_model_paths_table.py index 61710f64..5641e488 100644 --- a/migrations/versions/4e0c3d195a49_add_model_paths_table.py +++ b/migrations/versions/4e0c3d195a49_add_model_paths_table.py @@ -1,7 +1,7 @@ """add model paths table Revision ID: 4e0c3d195a49 -Revises: 7f2843d3f4e4 +Revises: 9c4d8e2f1a6b Create Date: 2026-07-24 21:14:39.062179 """ @@ -11,19 +11,19 @@ from alembic import op # revision identifiers, used by Alembic. revision = "4e0c3d195a49" -down_revision = "7f2843d3f4e4" +down_revision = "9c4d8e2f1a6b" branch_labels = None depends_on = None def upgrade() -> None: - # ### commands auto generated by Alembic - please adjust! ### op.create_table( "model_paths", sa.Column("id", sa.Integer(), nullable=False), sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False), sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.Column("updated_at", sa.Integer(), nullable=False, server_default="0"), sa.ForeignKeyConstraint( ["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE" ), @@ -35,21 +35,16 @@ def upgrade() -> None: name="uq_model_paths_model_path_provider", ), ) - op.create_index( - op.f("ix_model_paths_model_id"), "model_paths", ["model_id"], unique=False - ) + # No standalone index on model_id: the unique constraint's autoindex already + # leads on model_id. op.create_index( op.f("ix_model_paths_upstream_provider_id"), "model_paths", ["upstream_provider_id"], unique=False, ) - # ### end Alembic commands ### def downgrade() -> None: - # ### commands auto generated by Alembic - please adjust! ### op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths") - op.drop_index(op.f("ix_model_paths_model_id"), table_name="model_paths") op.drop_table("model_paths") - # ### end Alembic commands ### diff --git a/routstr/core/db.py b/routstr/core/db.py index 3cd34dcc..8b3586ec 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -222,7 +222,10 @@ async def release_stale_reservations( if released: logger.warning( "Released stale reservations", - extra={"released_reservations": released, "max_age_seconds": max_age_seconds}, + extra={ + "released_reservations": released, + "max_age_seconds": max_age_seconds, + }, ) return released @@ -255,9 +258,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in .where(col(ApiKey.total_spent) == 0) .where(col(ApiKey.total_requests) == 0) .where(col(ApiKey.parent_key_hash).is_(None)) - .where( - (col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff) - ) + .where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)) .where(~pending_invoice) .where(~has_children) ) @@ -331,8 +332,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore ), ) id: int | None = Field(default=None, primary_key=True) + # No standalone index on model_id: the unique constraint's autoindex already + # leads on model_id, so a second index only adds write amplification. model_id: str = Field( - index=True, description="Client-visible /v1/models id (forwarded_model_id or id)" + description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( description="Provider path stamped on chat completion responses, e.g. " @@ -344,6 +347,10 @@ class ModelPathRow(SQLModel, table=True): # type: ignore ondelete="CASCADE", description="upstream_providers.id this path was discovered from", ) + updated_at: int = Field( + default=0, + description="Unix timestamp of the refresh cycle that wrote this row", + ) class LightningInvoice(SQLModel, table=True): # type: ignore @@ -631,9 +638,7 @@ class CliToken(SQLModel, table=True): # type: ignore """Long-lived authorization token for CLI/agent use against admin endpoints.""" __tablename__ = "cli_tokens" - id: str = Field( - primary_key=True, default_factory=lambda: uuid.uuid4().hex - ) + id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex) token: str = Field(unique=True, index=True, description="Bearer token value") name: str = Field(description="Human-readable label for this token") created_at: int = Field(default_factory=lambda: int(time.time())) @@ -732,9 +737,7 @@ async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool: return result.rowcount == 1 -async def complete_routstr_fee_payout( - session: AsyncSession, paid_msats: int -) -> bool: +async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) -> bool: """Mark a checkpointed payout complete after the external payment succeeds.""" stmt = ( update(RoutstrFee) diff --git a/routstr/core/main.py b/routstr/core/main.py index 6eca6b52..979f5cdb 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -131,12 +131,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: refresh_upstreams_models_periodically(get_upstreams) ) model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) - if global_settings.model_paths_refresh_interval_seconds > 0: - from ..upstream.model_paths import refresh_model_paths_periodically + # Always started: the loop re-reads the enable flag and interval every + # iteration, so 0 -> N (or re-enabling) takes effect without a restart. + from ..upstream.model_paths import refresh_model_paths_periodically - model_paths_refresh_task = asyncio.create_task( - refresh_model_paths_periodically(get_upstreams) - ) + model_paths_refresh_task = asyncio.create_task( + refresh_model_paths_periodically(get_upstreams) + ) payout_task = asyncio.create_task(periodic_payout()) if global_settings.nsec: nip91_task = asyncio.create_task(announce_provider()) @@ -144,9 +145,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: if global_settings.providers_refresh_interval_seconds > 0: providers_task = asyncio.create_task(providers_cache_refresher()) key_reset_task = asyncio.create_task(periodic_key_reset()) - stale_reservation_task = asyncio.create_task( - periodic_stale_reservation_sweep() - ) + stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep()) dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) @@ -256,9 +255,7 @@ class _ImmutableStaticFiles(StaticFiles): async def get_response(self, path: str, scope: Scope) -> StarletteResponse: response = await super().get_response(path, scope) if response.status_code == 200: - response.headers["Cache-Control"] = ( - "public, max-age=31536000, immutable" - ) + response.headers["Cache-Control"] = "public, max-age=31536000, immutable" return response @@ -332,9 +329,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): # Serve the App Router RSC payload for the home page. @app.get("/index.txt", include_in_schema=False) async def serve_root_rsc() -> FileResponse: - return FileResponse( - UI_DIST_PATH / "index.txt", media_type="text/x-component" - ) + return FileResponse(UI_DIST_PATH / "index.txt", media_type="text/x-component") # Next.js is built with `trailingSlash: true`, so all UI page URLs end # with a slash (e.g. `/login/`). The proxy router catches `/{path:path}` diff --git a/routstr/core/settings.py b/routstr/core/settings.py index f5d0d2fd..ab91c3a4 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -100,8 +100,13 @@ class Settings(BaseSettings): ) enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") + enable_model_paths_refresh: bool = Field( + default=True, env="ENABLE_MODEL_PATHS_REFRESH" + ) refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") - refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS") + refund_sweep_ttl_seconds: int = Field( + default=604800, env="REFUND_SWEEP_TTL_SECONDS" + ) # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") @@ -120,9 +125,8 @@ class Settings(BaseSettings): # Discovery relays: list[str] = Field(default_factory=list, env="RELAYS") - enable_analytics_sharing: bool = Field( - default=True, env="ENABLE_ANALYTICS_SHARING" - ) + enable_analytics_sharing: bool = Field(default=True, env="ENABLE_ANALYTICS_SHARING") + def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]: """Discard unknown keys from persisted settings.""" @@ -330,7 +334,11 @@ class SettingsService: valid_fields = set(env_resolved.dict().keys()) merged_dict: dict[str, Any] = dict(env_resolved.dict()) merged_dict.update( - {k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields} + { + k: v + for k, v in db_json.items() + if v not in (None, "", [], {}) and k in valid_fields + } ) merged_dict = Settings(**merged_dict).dict() diff --git a/routstr/payment/models.py b/routstr/payment/models.py index b1e018d2..c433ddfa 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -455,7 +455,9 @@ async def _update_sats_pricing_once() -> None: for m in upstream.get_cached_models() ] upstream._models_cache = updated_models - upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models} + upstream._models_by_id = { + m.forwarded_model_id or m.id: m for m in updated_models + } updated_count += len(updated_models) if updated_count > 0: @@ -510,9 +512,7 @@ class ModelTestRequest(V2BaseModel): request_data: dict -@models_router.post( - "/api/models/test", dependencies=[Depends(_require_admin_api)] -) +@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) async def test_model( payload: ModelTestRequest, session: AsyncSession = Depends(get_session), @@ -601,7 +601,7 @@ async def model_paths() -> dict: """All models with every upstream provider path they are reachable through.""" from ..upstream.model_paths import get_all_model_paths - return {"data": await get_all_model_paths()} + return await get_all_model_paths() @models_router.get("/v1/models/paths/model") @@ -615,7 +615,7 @@ async def model_paths_for_model(model_id: str) -> dict: """ from ..upstream.model_paths import get_paths_for_model - return {"data": await get_paths_for_model(model_id)} + return await get_paths_for_model(model_id) @models_router.get("/v1/models") diff --git a/routstr/proxy.py b/routstr/proxy.py index c0b8794f..d3e9b59f 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -183,6 +183,19 @@ async def refresh_model_maps() -> None: disabled_model_keys=disabled_model_keys, ) + # Keep model-path discovery in sync with admin mutations: disabling or + # deleting a provider must stop advertising its paths immediately rather + # than after the next timed refresh. + from .upstream.model_paths import prune_model_paths_for_inactive_providers + + try: + await prune_model_paths_for_inactive_providers() + except Exception as e: # noqa: BLE001 - discovery sync must not break routing + logger.warning( + "Failed to prune model paths for inactive providers", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + async def refresh_model_maps_periodically() -> None: """Background task to refresh model maps every minute.""" diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index a8dba7e3..5f9116c9 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -237,6 +237,28 @@ class BaseUpstreamProvider: except (TypeError, ValueError): pass + def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: + """Discovery path for a reported sub-provider name. + + Must produce exactly the value ``_apply_provider_field`` would stamp on + a response whose upstream payload reported ``sub_provider``, so the + model-path discovery API never advertises a path that cannot appear on + the wire. Subclasses that override ``_apply_provider_field`` must + override this to match. + """ + provider_type = (self.provider_type or "").strip() + if not provider_type: + return None + sub = (sub_provider or "").strip() + if not sub or sub == provider_type or sub.startswith(f"{provider_type}:"): + return sub or provider_type + return f"{provider_type}:{sub}" + + def discovery_base_paths(self) -> list[str]: + """Paths stamped when the upstream reports no sub-provider of its own.""" + provider_type = (self.provider_type or "").strip() + return [provider_type] if provider_type else [] + def _apply_provider_field(self, response_json: object) -> None: """Stamp the routstr ``provider`` field onto an upstream response payload. diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 94d21833..beaafa6a 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -5,32 +5,34 @@ This is discovery/visibility data only — routing still selects the cheapest or best provider separately. A *path* is the provider string that may appear in Routstr chat completion -responses (see ``BaseUpstreamProvider._apply_provider_field``): +responses. The strings emitted here are produced by the provider's own +``discovery_path_for_subprovider`` / ``discovery_base_paths`` hooks, which +mirror ``_apply_provider_field`` so discovery and response stamping cannot +drift: - Direct upstream -> ```` e.g. ``anthropic`` - Generic/custom OpenRouter-compatible upstream -> ``generic:`` - Native OpenRouter routing to a sub-provider -> ``openrouter:`` - -Native OpenRouter does not emit a useful bare ``openrouter`` path when no -sub-provider is present; it reports ``unknown`` instead. +- Native OpenRouter with no usable sub-provider -> ``unknown`` """ from __future__ import annotations import asyncio import random +import time from typing import TYPE_CHECKING, Callable import httpx +from sqlalchemy import insert, or_ from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session from ..core.logging import get_logger -from .base import BaseUpstreamProvider if TYPE_CHECKING: - from ..payment.models import Model + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -39,6 +41,20 @@ logger = get_logger(__name__) _OPENROUTER_CONCURRENCY = 5 _OPENROUTER_TIMEOUT_SECONDS = 10.0 +# Rows inserted per statement during persist. Keeps each INSERT bounded while +# avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle. +_PERSIST_CHUNK_SIZE = 500 + +# Visibility key used across this module: routing carries the provider +# dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)), +# so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps. +ModelKey = tuple[str, int] + + +def _make_http_client() -> httpx.AsyncClient: + """Client factory, separated so tests can substitute a mock transport.""" + return httpx.AsyncClient() + def is_openrouter_base_url(base_url: str | None) -> bool: """True when ``base_url`` points at OpenRouter. @@ -61,19 +77,21 @@ def exposed_model_id(model: object) -> str: def public_model_id(model_id: str) -> str: """Model id exposed by model-path API responses. - Provider-prefixed ids such as ``z-ai/glm-5v-turbo`` are returned as - ``glm-5v-turbo`` so clients can search and display the same unqualified id - they pass to ``/v1/models/paths/model``. + Uses the same rule as ``create_model_mappings.get_base_model_id`` and + ``resolve_model_alias`` — strip everything before the *first* slash — so + the id shown here can be sent back to ``/v1/chat/completions`` verbatim. """ - return model_id.rsplit("/", 1)[-1] + return model_id.split("/", 1)[1] if "/" in model_id else model_id def openrouter_author_slug(model: object) -> str | None: """Return a canonical ``author/slug`` for the OpenRouter endpoints API. - OpenRouter requires the canonical id, never ``forwarded_model_id``. Prefer - ``canonical_slug``, then a slash-containing ``id``; otherwise there is no - usable form and endpoint discovery is skipped for this model. + Prefer ``canonical_slug``, then a slash-containing ``id``, then a + slash-containing ``forwarded_model_id``. The forwarded id is exactly what + the proxy sends upstream for admin-created alias rows (``base.py`` forwards + ``forwarded_model_id or id``), so it is a valid OpenRouter id when the + bare ``id`` is a local alias with no slash. """ canonical = getattr(model, "canonical_slug", None) if canonical and "/" in canonical: @@ -81,24 +99,51 @@ def openrouter_author_slug(model: object) -> str | None: model_id = getattr(model, "id", None) if model_id and "/" in model_id: return model_id + forwarded = getattr(model, "forwarded_model_id", None) + if forwarded and "/" in forwarded: + return forwarded return None -async def _fetch_openrouter_endpoint_paths( +class _RefreshCycleState: + """Per-refresh shared state: fetch dedupe cache and rate-limit latch. + + ``endpoint_cache`` dedupes byte-identical ``/endpoints`` fetches when two + providers point at the same OpenRouter base URL. ``rate_limited`` latches + on the first 429 so the rest of the cycle stops hammering a throttled API; + the whole provider result then degrades to "unknown" instead of an empty + list, which preserves previously persisted rows. + """ + + def __init__(self) -> None: + self.endpoint_cache: dict[tuple[str, str], list[str] | None] = {} + self.rate_limited = False + + +async def _fetch_openrouter_endpoint_subproviders( client: httpx.AsyncClient, base_url: str, api_key: str, author_slug: str, - path_prefix: str, semaphore: asyncio.Semaphore, -) -> list[str]: - """Return ``:`` paths for one model, or ``[]``. + cycle: _RefreshCycleState, +) -> list[str] | None: + """Return sub-provider names for one model, or ``None`` when unknown. - Failures (network, rate limit, bad payload) are logged and swallowed so one - model never breaks the whole refresh. + ``None`` (not ``[]``) signals a degraded fetch — network failure, rate + limit, non-200, or an unparseable payload — so callers can distinguish + "this model has no endpoints" from "we could not find out". Failures are + logged and swallowed so one model never breaks the whole refresh. """ + cache_key = (base_url, author_slug) + if cache_key in cycle.endpoint_cache: + return cycle.endpoint_cache[cache_key] + if cycle.rate_limited: + return None + url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} + result: list[str] | None async with semaphore: try: resp = await client.get( @@ -109,48 +154,60 @@ async def _fetch_openrouter_endpoint_paths( "OpenRouter endpoint discovery request failed", extra={"author_slug": author_slug, "error": str(e)}, ) - return [] + cycle.endpoint_cache[cache_key] = None + return None if resp.status_code == 429: logger.warning( - "OpenRouter endpoint discovery rate-limited", + "OpenRouter endpoint discovery rate-limited; aborting cycle", extra={"author_slug": author_slug}, ) - return [] + cycle.rate_limited = True + cycle.endpoint_cache[cache_key] = None + return None if resp.status_code != 200: logger.warning( "OpenRouter endpoint discovery non-200", extra={"author_slug": author_slug, "status_code": resp.status_code}, ) - return [] + cycle.endpoint_cache[cache_key] = None + return None try: endpoints = resp.json().get("data", {}).get("endpoints", []) + if not isinstance(endpoints, list): + endpoints = [] + names: list[str] = [] + for endpoint in endpoints: + provider_name = ( + endpoint.get("provider_name") if isinstance(endpoint, dict) else None + ) + if provider_name: + names.append(provider_name) + result = list(dict.fromkeys(names)) except Exception as e: # noqa: BLE001 logger.warning( "OpenRouter endpoint discovery bad payload", extra={"author_slug": author_slug, "error": str(e)}, ) - return [] + result = None - paths: list[str] = [] - for endpoint in endpoints: - provider_name = (endpoint or {}).get("provider_name") - if provider_name: - paths.append(f"{path_prefix}:{provider_name}") - # De-duplicate while preserving order. - return list(dict.fromkeys(paths)) + cycle.endpoint_cache[cache_key] = result + return result async def _load_model_visibility() -> tuple[ - dict[str, tuple[ModelRow, float]], set[str], set[int] + dict[ModelKey, ModelRow], set[ModelKey], set[int] ]: """Load the same DB model visibility inputs used by routing. ``refresh_model_maps`` builds routing from enabled providers, enabled DB - override rows, and disabled model ids. Model-path discovery uses the same - view so the discovery API does not advertise models routing would hide and - reports forwarded aliases from DB overrides consistently with ``/v1/models``. + override rows, and disabled model keys — all keyed on + ``(model_id.lower(), upstream_provider_id)`` because ``ModelRow``'s primary + key is composite and the same id legitimately exists on several providers. + Model-path discovery uses the same keying so disabling a model on one + provider never hides it on another, and one provider's + ``forwarded_model_id`` alias is never applied to a different provider. """ async with create_session() as session: query = select(UpstreamProviderRow).options( @@ -158,149 +215,149 @@ async def _load_model_visibility() -> tuple[ ) provider_rows = (await session.exec(query)).all() - overrides_by_id: dict[str, tuple[ModelRow, float]] = {} - disabled_model_ids: set[str] = set() + overrides_by_key: dict[ModelKey, ModelRow] = {} + disabled_model_keys: set[ModelKey] = set() enabled_provider_ids: set[int] = set() for provider in provider_rows: - if not provider.enabled: + if not provider.enabled or provider.id is None: continue - if provider.id is not None: - enabled_provider_ids.add(provider.id) + enabled_provider_ids.add(provider.id) for model in provider.models: + key = (model.id.lower(), provider.id) if model.enabled: - overrides_by_id[model.id] = (model, provider.provider_fee) + overrides_by_key[key] = model else: - disabled_model_ids.add(model.id) + disabled_model_keys.add(key) - return overrides_by_id, disabled_model_ids, enabled_provider_ids - - -def _row_to_visible_model( - model_id: str, - row: ModelRow, - provider_fee: float, -) -> Model | None: - """Convert an enabled DB override row into a routed model object.""" - from ..payment.models import _row_to_model - - try: - return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee) - except Exception as exc: # noqa: BLE001 - skip invalid override row - logger.warning( - "Skipping invalid model override while collecting model paths", - extra={ - "model_id": model_id, - "upstream_provider_id": getattr(row, "upstream_provider_id", None), - "error": str(exc), - "error_type": type(exc).__name__, - }, - ) - return None + return overrides_by_key, disabled_model_keys, enabled_provider_ids def _apply_model_visibility( upstream: BaseUpstreamProvider, - overrides_by_id: dict[str, tuple[ModelRow, float]] | None, - disabled_model_ids: set[str] | None, + overrides_by_key: dict[ModelKey, ModelRow] | None, + disabled_model_keys: set[ModelKey] | None, ) -> list[object]: - """Return provider models after DB disabled/override state is applied.""" - overrides_by_id = overrides_by_id or {} - disabled_model_ids = disabled_model_ids or set() + """Return provider models after DB disabled/override state is applied. + + Only the identity fields (``id``, ``forwarded_model_id``, + ``canonical_slug``) matter for path discovery, so DB override rows are used + directly rather than rebuilt into fully priced ``Model`` objects — the + pricing pipeline costs ~0.7ms of event-loop CPU per row for data this + module immediately discards. + """ + overrides_by_key = overrides_by_key or {} + disabled_model_keys = disabled_model_keys or set() + upstream_provider_id = getattr(upstream, "db_id", None) + if not isinstance(upstream_provider_id, int): + return [ + model + for model in upstream.get_cached_models() + if getattr(model, "enabled", True) + ] + visible_models: list[object] = [] seen_model_ids: set[str] = set() for model in upstream.get_cached_models(): model_id = getattr(model, "id", "") - if not getattr(model, "enabled", True) or model_id in disabled_model_ids: + key = (model_id.lower(), upstream_provider_id) + if not getattr(model, "enabled", True) or key in disabled_model_keys: continue - - if model_id in overrides_by_id: - override_row, provider_fee = overrides_by_id[model_id] - visible_model = _row_to_visible_model(model_id, override_row, provider_fee) - if visible_model is None: - continue - model = visible_model - - if not getattr(model, "enabled", True): - continue - visible_models.append(model) + # Apply overrides only for this provider's own model row. + override_row = overrides_by_key.get(key) + visible: object = model if override_row is None else override_row + visible_models.append(visible) seen_model_ids.add(model_id.lower()) - upstream_provider_id = getattr(upstream, "db_id", None) - if isinstance(upstream_provider_id, int): - for model_id, (override_row, provider_fee) in overrides_by_id.items(): - if model_id in disabled_model_ids: - continue - if ( - getattr(override_row, "upstream_provider_id", None) - != upstream_provider_id - ): - continue - if model_id.lower() in seen_model_ids: - continue - override_model = _row_to_visible_model(model_id, override_row, provider_fee) - if override_model is None: - continue - if getattr(override_model, "enabled", True): - visible_models.append(override_model) - seen_model_ids.add(model_id.lower()) + # DB-only override rows for this provider with no cached counterpart. + for (model_id_lower, provider_id), override_row in overrides_by_key.items(): + if provider_id != upstream_provider_id: + continue + if model_id_lower in seen_model_ids: + continue + visible_models.append(override_row) + seen_model_ids.add(model_id_lower) return visible_models async def _collect_provider_paths( upstream: BaseUpstreamProvider, - overrides_by_id: dict[str, tuple[ModelRow, float]] | None = None, - disabled_model_ids: set[str] | None = None, -) -> list[tuple[str, str]]: + overrides_by_key: dict[ModelKey, ModelRow] | None = None, + disabled_model_keys: set[ModelKey] | None = None, + cycle: _RefreshCycleState | None = None, +) -> list[tuple[str, str]] | None: """Collect ``(model_id, path)`` pairs for one provider instance. - Emits the direct ```` path for normal upstreams. For - OpenRouter-compatible providers, emits one path per OpenRouter sub-provider - endpoint, prefixed the same way response stamping prefixes it. + Emits the provider's ``discovery_base_paths`` for normal upstreams. For + OpenRouter-compatible providers, additionally emits one path per OpenRouter + sub-provider endpoint via ``discovery_path_for_subprovider`` so the strings + match response stamping exactly. + + Returns ``None`` when the provider's path set could not be determined this + cycle (every endpoint fetch degraded); callers must then keep previously + persisted rows instead of wiping them. """ - provider_type = (upstream.provider_type or "").strip() - models = _apply_model_visibility(upstream, overrides_by_id, disabled_model_ids) - is_openrouter = is_openrouter_base_url(upstream.base_url) + cycle = cycle or _RefreshCycleState() + models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) + base_paths = upstream.discovery_base_paths() - pairs: list[tuple[str, str]] = [] - if not is_openrouter: - for model in models: - if provider_type: - pairs.append((exposed_model_id(model), provider_type)) - return pairs + if not is_openrouter_base_url(upstream.base_url): + return [ + (exposed_model_id(model), path) for model in models for path in base_paths + ] - if not provider_type: - return pairs + if not (upstream.provider_type or "").strip(): + return [] + any_fetch_succeeded = False + any_fetch_attempted = False semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) - async with httpx.AsyncClient() as client: + async with _make_http_client() as client: async def _for_model(model: object) -> list[tuple[str, str]]: + nonlocal any_fetch_succeeded, any_fetch_attempted + model_id = exposed_model_id(model) + # Base paths always apply: responses whose upstream payload lacks a + # provider field are stamped with them (see _apply_provider_field). + pairs = [(model_id, path) for path in base_paths] author_slug = openrouter_author_slug(model) if not author_slug: - return [] - paths = await _fetch_openrouter_endpoint_paths( + return pairs + any_fetch_attempted = True + sub_providers = await _fetch_openrouter_endpoint_subproviders( client, upstream.base_url, upstream.api_key, author_slug, - provider_type, semaphore, + cycle, ) - model_id = exposed_model_id(model) - return [(model_id, path) for path in paths] + if sub_providers is None: + return [] + any_fetch_succeeded = True + paths = [ + upstream.discovery_path_for_subprovider(name) for name in sub_providers + ] + pairs.extend((model_id, path) for path in paths if path) + return list(dict.fromkeys(pairs)) results = await asyncio.gather( *(_for_model(m) for m in models), return_exceptions=True ) + if any_fetch_attempted and not any_fetch_succeeded: + # Every endpoint lookup degraded (offline, throttled, bad payloads): + # the true path set is unknown, not empty. + return None + + pairs: list[tuple[str, str]] = [] for result in results: if isinstance(result, BaseException): logger.warning( "OpenRouter endpoint discovery task errored", - extra={"provider": provider_type, "error": str(result)}, + extra={"provider": upstream.provider_type, "error": str(result)}, ) continue pairs.extend(result) @@ -314,33 +371,57 @@ async def _persist_provider_paths( """Replace all rows for ``upstream_provider_id`` with ``pairs``. Replacement (not upsert) so stale paths disappear when provider config or - upstream availability changes. + upstream availability changes. Rows are written with chunked bulk INSERTs + so the transaction holds SQLite's write lock briefly — billing writes share + this database file. """ unique_pairs = list(dict.fromkeys(pairs)) + now = int(time.time()) async with create_session() as session: await session.exec( # type: ignore[call-overload] delete(ModelPathRow).where( col(ModelPathRow.upstream_provider_id) == upstream_provider_id ) ) - for model_id, path in unique_pairs: - session.add( - ModelPathRow( - model_id=model_id, - path=path, - upstream_provider_id=upstream_provider_id, - ) + for start in range(0, len(unique_pairs), _PERSIST_CHUNK_SIZE): + chunk = unique_pairs[start : start + _PERSIST_CHUNK_SIZE] + await session.execute( + insert(ModelPathRow), + [ + { + "model_id": model_id, + "path": path, + "upstream_provider_id": upstream_provider_id, + "updated_at": now, + } + for model_id, path in chunk + ], ) await session.commit() -async def _prune_inactive_provider_paths(active_provider_ids: set[int]) -> None: - """Delete paths for providers no longer present in the live upstream set.""" +async def prune_model_paths_for_inactive_providers() -> None: + """Delete paths whose provider is no longer enabled in the database. + + Called from ``refresh_model_maps`` so admin mutations (disable/delete + provider) stop advertising a provider's paths immediately instead of + waiting for the next timed refresh. Uses the DB as the source of truth, so + it is safe at boot even before upstreams initialize. + """ async with create_session() as session: + enabled_ids = ( + await session.exec( + select(UpstreamProviderRow.id).where( + col(UpstreamProviderRow.enabled).is_(True) + ) + ) + ).all() stmt = delete(ModelPathRow) - if active_provider_ids: + if enabled_ids: stmt = stmt.where( - col(ModelPathRow.upstream_provider_id).not_in(active_provider_ids) + col(ModelPathRow.upstream_provider_id).not_in( + [pid for pid in enabled_ids if pid is not None] + ) ) await session.exec(stmt) # type: ignore[call-overload] await session.commit() @@ -352,28 +433,42 @@ async def refresh_model_paths( """Recompute and persist model paths for every enabled provider. One provider's failure is logged and isolated; it must not break the rest. + A provider whose paths could not be determined this cycle keeps its + previously persisted rows. An empty ``upstreams`` list (e.g. a failed + ``initialize_upstreams`` at boot) is treated as "unknown" and touches + nothing. """ + if not upstreams: + logger.warning("Skipping model paths refresh: no live upstreams") + return + ( - overrides_by_id, - disabled_model_ids, + overrides_by_key, + disabled_model_keys, enabled_provider_ids, ) = await _load_model_visibility() - active_provider_ids = { - upstream.db_id - for upstream in upstreams - if upstream.db_id is not None and upstream.db_id in enabled_provider_ids - } - await _prune_inactive_provider_paths(active_provider_ids) + await prune_model_paths_for_inactive_providers() + cycle = _RefreshCycleState() for upstream in upstreams: if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: continue try: pairs = await _collect_provider_paths( upstream, - overrides_by_id=overrides_by_id, - disabled_model_ids=disabled_model_ids, + overrides_by_key=overrides_by_key, + disabled_model_keys=disabled_model_keys, + cycle=cycle, ) + if pairs is None: + logger.warning( + "Model paths unknown this cycle; keeping previous rows", + extra={ + "provider": upstream.provider_type or upstream.base_url, + "db_id": upstream.db_id, + }, + ) + continue await _persist_provider_paths(upstream.db_id, pairs) except Exception as e: # noqa: BLE001 - isolate per-provider failures logger.error( @@ -387,18 +482,27 @@ async def refresh_model_paths( ) +def _refresh_interval_seconds() -> int: + """Current interval, re-read every loop so runtime setting changes apply.""" + from ..core.settings import settings + + if not getattr(settings, "enable_model_paths_refresh", True): + return 0 + return int(getattr(settings, "model_paths_refresh_interval_seconds", 0) or 0) + + async def refresh_model_paths_periodically( upstreams_provider: ( Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider] ), ) -> None: - """Background task mirroring ``refresh_upstreams_models_periodically``.""" - from ..core.settings import settings + """Background task mirroring ``refresh_upstreams_models_periodically``. - interval = getattr(settings, "model_paths_refresh_interval_seconds", 0) - if not interval or interval <= 0: - logger.info("Model paths refresh disabled (interval <= 0)") - return + The interval and enable flag are re-read every iteration, so the refresh + can be turned off (or on) and retuned without a restart. While disabled the + task idles instead of exiting, so re-enabling takes effect. + """ + _DISABLED_POLL_SECONDS = 60.0 def _resolve_upstreams() -> list[BaseUpstreamProvider]: if callable(upstreams_provider): @@ -406,6 +510,14 @@ async def refresh_model_paths_periodically( return upstreams_provider while True: + interval = _refresh_interval_seconds() + if interval <= 0: + try: + await asyncio.sleep(_DISABLED_POLL_SECONDS) + except asyncio.CancelledError: + break + continue + try: await refresh_model_paths(_resolve_upstreams()) except asyncio.CancelledError: @@ -423,50 +535,79 @@ async def refresh_model_paths_periodically( break -async def get_all_model_paths() -> list[dict]: +async def get_all_model_paths() -> dict: """All models with their paths, shaped for ``GET /v1/models/paths``.""" async with create_session() as session: rows = ( - await session.exec(select(ModelPathRow).order_by(ModelPathRow.model_id)) + await session.exec( + select(ModelPathRow).order_by( + col(ModelPathRow.model_id), + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) + ) ).all() grouped: dict[str, list[dict]] = {} seen_paths: dict[str, set[str]] = {} + updated_at = 0 for row in rows: + updated_at = max(updated_at, row.updated_at) model_id = public_model_id(row.model_id) if row.path in seen_paths.setdefault(model_id, set()): continue seen_paths[model_id].add(row.path) grouped.setdefault(model_id, []).append({"path": row.path}) - return [{"id": model_id, "paths": paths} for model_id, paths in grouped.items()] + # Deterministic output: models sorted by public id, paths sorted within. + data: list[dict] = [] + for grouped_model_id in sorted(grouped): + model_paths = sorted(grouped[grouped_model_id], key=lambda p: str(p["path"])) + data.append({"id": grouped_model_id, "paths": model_paths}) + return {"data": data, "updated_at": updated_at or None} -async def get_paths_for_model(model_id: str) -> list[dict]: +async def get_paths_for_model(model_id: str) -> dict: """Paths for a single model, shaped for ``GET /v1/models/paths/model``. Match by the public, unqualified model id, mirroring the model cache alias behavior. Both ``deepseek-v4-pro`` and ``deepseek/deepseek-v4-pro`` resolve - every row whose stored id has the same base model id. + every row whose stored id has the same base model id. The candidate set is + narrowed in SQL (exact id or ``%/`` suffix) so the route does not + materialize the whole table per request. """ - requested_id = public_model_id(model_id) + # The request may be a full stored id ("z-ai/glm-5v-turbo") or an + # already-stripped public id ("fireworks/models/glm-5"); accept both. + accepted_ids = {model_id, public_model_id(model_id)} async with create_session() as session: + conditions = [] + for candidate in accepted_ids: + conditions.append(col(ModelPathRow.model_id) == candidate) + conditions.append(col(ModelPathRow.model_id).endswith(f"/{candidate}")) rows = ( await session.exec( - select(ModelPathRow).order_by( - ModelPathRow.path, + select(ModelPathRow) + .where(or_(*conditions)) + .order_by( + col(ModelPathRow.path), col(ModelPathRow.upstream_provider_id), - ModelPathRow.model_id, + col(ModelPathRow.model_id), ) ) ).all() seen: set[str] = set() paths: list[dict] = [] + updated_at = 0 for row in rows: - if public_model_id(row.model_id) != requested_id: + # The SQL suffix match is a prefilter; enforce the exact public-id rule. + if ( + row.model_id not in accepted_ids + and public_model_id(row.model_id) not in accepted_ids + ): continue + updated_at = max(updated_at, row.updated_at) if row.path in seen: continue seen.add(row.path) paths.append({"path": row.path}) - return paths + return {"data": paths, "updated_at": updated_at or None} diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 1caeaa5c..fe92c4f9 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -18,6 +18,23 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): supports_anthropic_messages = True litellm_provider_prefix = "openrouter/" + def discovery_path_for_subprovider(self, sub_provider: str | None) -> str | None: + """Mirror ``_apply_provider_field``: strip repeated prefixes, map a + missing or self-echoing sub-provider to the literal ``"unknown"``.""" + provider_type = (self.provider_type or "").strip() + sub = (sub_provider or "").strip() + prefix = f"{provider_type}:" + while sub.lower().startswith(prefix.lower()): + sub = sub[len(prefix) :].strip() + if not sub or sub.lower() == provider_type.lower(): + return "unknown" + return f"{provider_type}:{sub}" + + def discovery_base_paths(self) -> list[str]: + """Native OpenRouter never stamps a bare ``openrouter``; a response + with no sub-provider is stamped ``unknown``.""" + return ["unknown"] + def _apply_provider_field(self, response_json: object) -> None: """Stamp the ``provider`` field for OpenRouter responses. diff --git a/tests/unit/test_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index d820a056..c1ce13c8 100644 --- a/tests/unit/test_fee_payout_migration.py +++ b/tests/unit/test_fee_payout_migration.py @@ -37,7 +37,9 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None: "payout_in_progress_msats, payout_started_at FROM routstr_fees" ).fetchone() - assert version == ("9c4d8e2f1a6b",) + # Head of the 7f2843d3f4e4 lineage: model-paths chains onto the fee-payout + # repair migration. + assert version == ("4e0c3d195a49",) assert { "id", "accumulated_msats", diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 947c4312..b1747436 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -1,4 +1,10 @@ -"""Tests for the model-path discovery service and endpoints.""" +"""Tests for the model-path discovery service and endpoints. + +These tests exercise the public entry points (``refresh_model_paths``, +``get_all_model_paths``, ``get_paths_for_model``) rather than private helpers, +and fake OpenRouter at the transport level (``httpx.MockTransport``) so a +signature drift in the SUT fails loudly instead of silently returning ``[]``. +""" from __future__ import annotations @@ -7,11 +13,13 @@ import json import os from contextlib import asynccontextmanager from types import SimpleNamespace -from typing import Any, AsyncGenerator +from typing import Any, AsyncGenerator, Callable +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient +from sqlalchemy import event from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession @@ -22,6 +30,8 @@ os.environ.setdefault("UPSTREAM_API_KEY", "test") from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402 from routstr.payment.models import models_router # noqa: E402 from routstr.upstream import model_paths as mp # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 +from routstr.upstream.openrouter import OpenRouterUpstreamProvider # noqa: E402 # --------------------------------------------------------------------------- # # Fakes @@ -74,7 +84,10 @@ def _model_row( ) -class _FakeProvider: +class _FakeProvider(BaseUpstreamProvider): + """Real ``BaseUpstreamProvider`` so the discovery-path hooks are the + production ones, with cached models injected.""" + def __init__( self, *, @@ -84,37 +97,83 @@ class _FakeProvider: db_id: int | None = 1, api_key: str = "sk-test", ) -> None: - self.provider_type = provider_type - self.base_url = base_url - self.api_key = api_key + super().__init__(base_url=base_url, api_key=api_key) + self.provider_type = provider_type # shadow the class attribute self.db_id = db_id self._models = models - def get_cached_models(self) -> list[SimpleNamespace]: + def get_cached_models(self) -> list[SimpleNamespace]: # type: ignore[override] return self._models -class _FakeResponse: - def __init__(self, status_code: int, payload: Any) -> None: - self.status_code = status_code - self._payload = payload +class _FakeOpenRouterProvider(OpenRouterUpstreamProvider): + """Real OpenRouter provider so the ``unknown`` mapping is the production one.""" - def json(self) -> Any: - return self._payload + def __init__( + self, + *, + models: list[SimpleNamespace], + db_id: int | None = 2, + api_key: str = "sk-or", + ) -> None: + super().__init__(api_key=api_key) + self.db_id = db_id + self._models = models + + def get_cached_models(self) -> list[SimpleNamespace]: # type: ignore[override] + return self._models + + +def _mock_transport( + monkeypatch: pytest.MonkeyPatch, + handler: Callable[[httpx.Request], httpx.Response], +) -> dict[str, int]: + """Route the SUT's HTTP through ``httpx.MockTransport`` and count requests.""" + counter = {"requests": 0} + + def _counting_handler(request: httpx.Request) -> httpx.Response: + counter["requests"] += 1 + return handler(request) + + def _factory() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(_counting_handler)) + + monkeypatch.setattr(mp, "_make_http_client", _factory) + return counter + + +def _endpoints_response(*provider_names: str) -> httpx.Response: + return httpx.Response( + 200, + json={"data": {"endpoints": [{"provider_name": n} for n in provider_names]}}, + ) + + +_SEEDED_PROVIDER_IDS = (1, 2, 4, 5, 7) @pytest.fixture async def patched_session( monkeypatch: pytest.MonkeyPatch, ) -> AsyncGenerator[AsyncEngine, None]: - """Bind the service's ``create_session`` to a fresh in-memory engine.""" + """Bind the service's ``create_session`` to a fresh in-memory engine. + + Foreign keys are enforced (``PRAGMA foreign_keys=ON``) so a ModelPathRow + insert for an unseeded provider fails here even though production SQLite + currently runs with the pragma off. + """ engine = create_async_engine("sqlite+aiosqlite:///:memory:") + + @event.listens_for(engine.sync_engine, "connect") + def _enable_fk(dbapi_conn: Any, _record: Any) -> None: + dbapi_conn.execute("PRAGMA foreign_keys=ON") + async with engine.begin() as conn: await conn.run_sync(SQLModel.metadata.create_all) - # Seed the FK target so ModelPathRow inserts satisfy the constraint. + # Seed every provider id the tests insert path rows for. async with AsyncSession(engine) as session: - for pid in (1, 2): + for pid in _SEEDED_PROVIDER_IDS: session.add( UpstreamProviderRow( id=pid, @@ -136,6 +195,17 @@ async def patched_session( await engine.dispose() +def _paths_of(payload: dict, model_id: str) -> set[str]: + for entry in payload["data"]: + if entry["id"] == model_id: + return {p["path"] for p in entry["paths"]} + return set() + + +def _ids_of(payload: dict) -> set[str]: + return {entry["id"] for entry in payload["data"]} + + # --------------------------------------------------------------------------- # # Predicates / pure helpers # --------------------------------------------------------------------------- # @@ -161,12 +231,18 @@ def test_exposed_model_id_prefers_forwarded() -> None: assert mp.exposed_model_id(_model("claude-x")) == "claude-x" -def test_public_model_id_strips_provider_prefix() -> None: +def test_public_model_id_strips_first_provider_prefix() -> None: + """Must match ``create_model_mappings.get_base_model_id`` (first slash), + so the id shown by discovery can be sent to chat completions verbatim.""" assert mp.public_model_id("z-ai/glm-5v-turbo") == "glm-5v-turbo" assert mp.public_model_id("gpt-4o-mini") == "gpt-4o-mini" + assert ( + mp.public_model_id("accounts/fireworks/models/glm-5") + == "fireworks/models/glm-5" + ) -def test_openrouter_author_slug_uses_canonical_not_forwarded() -> None: +def test_openrouter_author_slug_prefers_canonical() -> None: m = _model( "claude-opus-4.6", forwarded_model_id="forwarded-only", @@ -180,216 +256,174 @@ def test_openrouter_author_slug_falls_back_to_slash_id() -> None: assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" +def test_openrouter_author_slug_falls_back_to_forwarded_id() -> None: + """Admin-created alias rows have a slash-less local id; the forwarded id is + what the proxy actually sends to OpenRouter, so it is a usable slug.""" + m = _model("my-alias", forwarded_model_id="anthropic/claude-opus-4.6") + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + def test_openrouter_author_slug_none_when_no_slash() -> None: m = _model("claude-opus-4.6", canonical_slug="claude-opus-4.6") assert mp.openrouter_author_slug(m) is None +def test_discovery_paths_mirror_response_stamping() -> None: + """The discovery hook and ``_apply_provider_field`` must agree.""" + generic = _FakeProvider(provider_type="generic", base_url="https://x", models=[]) + assert generic.discovery_path_for_subprovider("Anthropic") == "generic:Anthropic" + assert generic.discovery_base_paths() == ["generic"] + + native = _FakeOpenRouterProvider(models=[]) + assert native.discovery_path_for_subprovider("GMICloud") == "openrouter:GMICloud" + # Sub-provider echoing the router name is stamped "unknown" on responses. + assert native.discovery_path_for_subprovider("OpenRouter") == "unknown" + assert native.discovery_path_for_subprovider("openrouter:openrouter") == "unknown" + assert native.discovery_path_for_subprovider(None) == "unknown" + assert native.discovery_base_paths() == ["unknown"] + + # --------------------------------------------------------------------------- # -# Collection +# Refresh through the public entry point # --------------------------------------------------------------------------- # @pytest.mark.asyncio -async def test_direct_provider_single_path_uses_provider_type() -> None: +async def test_direct_provider_single_path_uses_provider_type( + patched_session: AsyncEngine, +) -> None: provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", models=[_model("claude-opus-4.6")], + db_id=1, ) - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("claude-opus-4.6", "anthropic")] + await mp.refresh_model_paths([provider]) + payload = await mp.get_all_model_paths() + assert payload["data"] == [ + {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + ] + assert payload["updated_at"] is not None @pytest.mark.asyncio -async def test_direct_path_stores_exposed_model_id() -> None: +async def test_direct_path_stores_exposed_model_id( + patched_session: AsyncEngine, +) -> None: provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", models=[_model("internal-id", forwarded_model_id="claude-opus-4.6")], + db_id=1, ) - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("claude-opus-4.6", "anthropic")] + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} @pytest.mark.asyncio -async def test_disabled_models_excluded() -> None: +async def test_disabled_cached_models_excluded( + patched_session: AsyncEngine, +) -> None: provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", - models=[ - _model("enabled-model"), - _model("disabled-model", enabled=False), - ], + models=[_model("enabled-model"), _model("disabled-model", enabled=False)], + db_id=1, ) - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("enabled-model", "anthropic")] + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"enabled-model"} @pytest.mark.asyncio -async def test_openrouter_provider_adds_endpoint_paths( - monkeypatch: pytest.MonkeyPatch, -) -> None: - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _fake_get( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - return _FakeResponse( - 200, - { - "data": { - "endpoints": [ - {"provider_name": "Anthropic"}, - {"provider_name": "Amazon Bedrock"}, - ] - } - }, - ) - - monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert ("claude-opus-4.6", "openrouter:Anthropic") in pairs - assert ("claude-opus-4.6", "openrouter:Amazon Bedrock") in pairs - assert ("claude-opus-4.6", "openrouter") not in pairs - assert len(pairs) == 2 - - -@pytest.mark.asyncio -async def test_generic_provider_with_openrouter_base_url_discovers( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """A generic provider pointed at OpenRouter exposes the response-stamped - ``generic:`` path, not a native ``openrouter:`` path.""" - provider = _FakeProvider( - provider_type="generic", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _fake_get( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - return _FakeResponse( - 200, {"data": {"endpoints": [{"provider_name": "Anthropic"}]}} - ) - - monkeypatch.setattr("httpx.AsyncClient.get", _fake_get) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [("claude-opus-4.6", "generic:Anthropic")] - - -@pytest.mark.asyncio -async def test_openrouter_failure_degrades_gracefully( - monkeypatch: pytest.MonkeyPatch, -) -> None: - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _boom( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - raise RuntimeError("network down") - - monkeypatch.setattr("httpx.AsyncClient.get", _boom) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [] - - -@pytest.mark.asyncio -async def test_openrouter_rate_limit_skips_model( - monkeypatch: pytest.MonkeyPatch, -) -> None: - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], - ) - - async def _rate_limited( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - return _FakeResponse(429, {}) - - monkeypatch.setattr("httpx.AsyncClient.get", _rate_limited) - - pairs = await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert pairs == [] - - -@pytest.mark.asyncio -async def test_openrouter_fanout_is_bounded( - monkeypatch: pytest.MonkeyPatch, -) -> None: - monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) - models = [ - _model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20) - ] - provider = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=models, - ) - - state = {"current": 0, "max": 0} - - async def _slow_get( - self: object, - url: str, - headers: object = None, - timeout: object = None, - ) -> _FakeResponse: - state["current"] += 1 - state["max"] = max(state["max"], state["current"]) - await asyncio.sleep(0.02) - state["current"] -= 1 - return _FakeResponse(200, {"data": {"endpoints": [{"provider_name": "X"}]}}) - - monkeypatch.setattr("httpx.AsyncClient.get", _slow_get) - - await mp._collect_provider_paths(provider) # type: ignore[arg-type] - assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" - - -# --------------------------------------------------------------------------- # -# Persistence + query -# --------------------------------------------------------------------------- # - - -@pytest.mark.asyncio -async def test_refresh_replaces_stale_rows( +async def test_disabling_model_on_one_provider_keeps_other_provider( patched_session: AsyncEngine, ) -> None: - await mp._persist_provider_paths(1, [("m1", "anthropic"), ("m2", "anthropic")]) - first = await mp.get_all_model_paths() - assert {row["id"] for row in first} == {"m1", "m2"} + """Regression for cross-provider isolation: ModelRow's primary key is + (id, upstream_provider_id), so a disable row on provider 2 must not hide + provider 1's model.""" + async with AsyncSession(patched_session) as session: + session.add(_model_row("shared-model", upstream_provider_id=2, enabled=False)) + await session.commit() - # Second refresh with a different set — stale m2 must disappear. - await mp._persist_provider_paths(1, [("m1", "anthropic")]) - second = await mp.get_all_model_paths() - assert {row["id"] for row in second} == {"m1"} + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("shared-model")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("shared-model")], + db_id=2, + ) + + await mp.refresh_model_paths([p1, p2]) + + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "shared-model") == {"anthropic"} + + +@pytest.mark.asyncio +async def test_override_alias_not_applied_across_providers( + patched_session: AsyncEngine, +) -> None: + """Provider 2's forwarded_model_id must never rename provider 1's model.""" + async with AsyncSession(patched_session) as session: + session.add( + _model_row( + "shared-model", + upstream_provider_id=2, + forwarded_model_id="private-alias", + ) + ) + await session.commit() + + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("shared-model")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("shared-model")], + db_id=2, + ) + + await mp.refresh_model_paths([p1, p2]) + + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "shared-model") == {"anthropic"} + assert _paths_of(payload, "private-alias") == {"generic"} + + +@pytest.mark.asyncio +async def test_override_matching_is_case_insensitive( + patched_session: AsyncEngine, +) -> None: + """Routing lowercases both sides when matching DB rows to cached models; + discovery must do the same for mixed-case ids.""" + async with AsyncSession(patched_session) as session: + session.add( + _model_row( + "deepseek-ai/deepseek-v4-flash", + upstream_provider_id=1, + forwarded_model_id="public-alias", + ) + ) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("deepseek-ai/DeepSeek-V4-Flash")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"public-alias"} @pytest.mark.asyncio @@ -407,9 +441,9 @@ async def test_refresh_model_paths_excludes_db_disabled_override( db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [] + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio @@ -427,9 +461,9 @@ async def test_refresh_model_paths_uses_db_forwarded_alias( db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [ + assert (await mp.get_all_model_paths())["data"] == [ {"id": "public-alias", "paths": [{"path": "anthropic"}]} ] @@ -449,44 +483,88 @@ async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cac db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [ + assert (await mp.get_all_model_paths())["data"] == [ {"id": "public-deployment", "paths": [{"path": "generic"}]} ] @pytest.mark.asyncio -async def test_refresh_model_paths_prunes_inactive_provider_rows( +async def test_refresh_replaces_stale_rows( patched_session: AsyncEngine, ) -> None: - await mp._persist_provider_paths(1, [("m1", "anthropic")]) - await mp._persist_provider_paths(2, [("m2", "openrouter:Anthropic")]) + p_two_models = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1"), _model("m2")], + db_id=1, + ) + await mp.refresh_model_paths([p_two_models]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1", "m2"} + p_one_model = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([p_one_model]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + +@pytest.mark.asyncio +async def test_refresh_with_no_upstreams_keeps_existing_rows( + patched_session: AsyncEngine, +) -> None: + """An empty live upstream list (e.g. failed boot init) means "unknown", + not "delete everything".""" provider = _FakeProvider( provider_type="anthropic", base_url="https://api.anthropic.com/v1", models=[_model("m1")], db_id=1, ) - - await mp.refresh_model_paths([provider]) # type: ignore[list-item] - active_only = await mp.get_all_model_paths() - assert active_only == [{"id": "m1", "paths": [{"path": "anthropic"}]}] + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} await mp.refresh_model_paths([]) - assert await mp.get_all_model_paths() == [] + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + +@pytest.mark.asyncio +async def test_prune_removes_rows_of_disabled_db_provider( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.enabled = False + session.add(provider_row) + await session.commit() + + await mp.prune_model_paths_for_inactive_providers() + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio async def test_refresh_model_paths_skips_disabled_db_provider( patched_session: AsyncEngine, ) -> None: - await mp._persist_provider_paths(1, [("stale-model", "anthropic")]) async with AsyncSession(patched_session) as session: provider_row = await session.get(UpstreamProviderRow, 1) assert provider_row is not None provider_row.enabled = False + session.add(provider_row) await session.commit() provider = _FakeProvider( @@ -496,109 +574,9 @@ async def test_refresh_model_paths_skips_disabled_db_provider( db_id=1, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] + await mp.refresh_model_paths([provider]) - assert await mp.get_all_model_paths() == [] - - -@pytest.mark.asyncio -async def test_same_model_two_providers_two_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) - - data = await mp.get_all_model_paths() - assert len(data) == 1 - entry = data[0] - assert entry["id"] == "claude-opus-4.6" - paths = {p["path"] for p in entry["paths"]} - assert paths == {"anthropic", "openrouter:Anthropic"} - # No canonical_id anywhere. - assert "canonical_id" not in entry - assert all("canonical_id" not in p for p in entry["paths"]) - - -@pytest.mark.asyncio -async def test_get_all_model_paths_deduplicates_visible_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) - - assert await mp.get_all_model_paths() == [ - {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} - ] - - -@pytest.mark.asyncio -async def test_get_all_model_paths_returns_unqualified_model_ids( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) - await mp._persist_provider_paths(5, [("openai/gpt-4o-mini", "openrouter:OpenAI")]) - - data = await mp.get_all_model_paths() - - assert {row["id"] for row in data} == {"glm-5v-turbo", "gpt-4o-mini"} - - -@pytest.mark.asyncio -async def test_get_paths_for_model_returns_only_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "openrouter:Anthropic")]) - - paths = await mp.get_paths_for_model("claude-opus-4.6") - assert {p["path"] for p in paths} == {"anthropic", "openrouter:Anthropic"} - assert all(set(p.keys()) == {"path"} for p in paths) - assert await mp.get_paths_for_model("does-not-exist") == [] - - -@pytest.mark.asyncio -async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(4, [("z-ai/glm-5v-turbo", "openrouter:Z.AI")]) - - paths = await mp.get_paths_for_model("glm-5v-turbo") - - assert paths == [{"path": "openrouter:Z.AI"}] - - -@pytest.mark.asyncio -async def test_get_paths_for_model_deduplicates_visible_paths( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(1, [("anthropic/claude-opus-4.6", "anthropic")]) - await mp._persist_provider_paths(2, [("claude-opus-4.6", "anthropic")]) - - assert await mp.get_paths_for_model("claude-opus-4.6") == [ - {"path": "anthropic"} - ] - assert await mp.get_paths_for_model("anthropic/claude-opus-4.6") == [ - {"path": "anthropic"} - ] - - -@pytest.mark.asyncio -async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( - patched_session: AsyncEngine, -) -> None: - await mp._persist_provider_paths(7, [("deepseek-v4-pro", "generic")]) - await mp._persist_provider_paths( - 4, [("deepseek/deepseek-v4-pro", "openrouter:DeepSeek")] - ) - - short_paths = await mp.get_paths_for_model("deepseek-v4-pro") - prefixed_paths = await mp.get_paths_for_model("deepseek/deepseek-v4-pro") - - assert short_paths == [ - {"path": "generic"}, - {"path": "openrouter:DeepSeek"}, - ] - assert prefixed_paths == short_paths + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio @@ -611,8 +589,8 @@ async def test_refresh_model_paths_skips_provider_without_db_id( models=[_model("claude-opus-4.6")], db_id=None, ) - await mp.refresh_model_paths([provider]) # type: ignore[list-item] - assert await mp.get_all_model_paths() == [] + await mp.refresh_model_paths([provider]) + assert (await mp.get_all_model_paths())["data"] == [] @pytest.mark.asyncio @@ -626,30 +604,476 @@ async def test_refresh_model_paths_isolates_provider_failure( db_id=1, ) bad = _FakeProvider( - provider_type="openrouter", - base_url="https://openrouter.ai/api/v1", - models=[_model("m", canonical_slug="a/m")], + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("m")], db_id=2, ) original = mp._collect_provider_paths - async def _maybe_fail( - upstream: Any, *args: Any, **kwargs: Any - ) -> list[tuple[str, str]]: + async def _maybe_fail(upstream: Any, *args: Any, **kwargs: Any) -> Any: if upstream is bad: raise RuntimeError("boom") - return await original(upstream, *args, **kwargs) # type: ignore[arg-type] + return await original(upstream, *args, **kwargs) monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) - await mp.refresh_model_paths([good, bad]) # type: ignore[list-item] - data = await mp.get_all_model_paths() - assert {row["id"] for row in data} == {"claude-opus-4.6"} + await mp.refresh_model_paths([good, bad]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} # --------------------------------------------------------------------------- # -# Endpoints +# OpenRouter endpoint discovery (transport-level fakes) +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_openrouter_provider_adds_endpoint_paths( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport( + monkeypatch, + lambda request: _endpoints_response("Anthropic", "Amazon Bedrock"), + ) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:Anthropic" in paths + assert "openrouter:Amazon Bedrock" in paths + assert "openrouter" not in paths + + +@pytest.mark.asyncio +async def test_openrouter_self_echoing_subprovider_maps_to_unknown( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Responses stamped for a sub-provider echoing "OpenRouter" say + ``unknown``; discovery must advertise the same string, never + ``openrouter:OpenRouter``.""" + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("OpenRouter")) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:OpenRouter" not in paths + assert "unknown" in paths + + +@pytest.mark.asyncio +async def test_generic_provider_with_openrouter_base_url_discovers( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A generic provider pointed at OpenRouter exposes both the bare + ``generic`` path (stamped when the upstream omits its provider field) and + the ``generic:`` endpoint paths.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=1, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert paths == {"generic", "generic:Anthropic"} + + +@pytest.mark.asyncio +async def test_openrouter_failure_keeps_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A transient upstream failure means the path set is unknown; previously + persisted rows must survive, mirroring ``refresh_models_cache``.""" + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:Anthropic" in before + + def _network_down(request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("network down", request=request) + + _mock_transport(monkeypatch, _network_down) + await mp.refresh_model_paths([provider]) + + after = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert after == before + + +@pytest.mark.asyncio +async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """The first 429 latches: no further endpoint requests this cycle, and the + provider's previously persisted rows survive.""" + models = [_model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(10)] + provider = _FakeOpenRouterProvider(models=models, db_id=2) + + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + assert _paths_of(await mp.get_all_model_paths(), "m0") == { + "unknown", + "openrouter:Anthropic", + } + + counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) + await mp.refresh_model_paths([provider]) + + # Up to _OPENROUTER_CONCURRENCY requests may already be in flight when the + # first 429 lands; the latch must stop everything after that. + assert counter["requests"] <= mp._OPENROUTER_CONCURRENCY, ( + "429 must abort the remaining fan-out" + ) + assert _paths_of(await mp.get_all_model_paths(), "m0") == { + "unknown", + "openrouter:Anthropic", + } + + +@pytest.mark.asyncio +async def test_openrouter_bad_payload_shapes_do_not_raise( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """``endpoints: null`` and non-list endpoint payloads are swallowed as + documented, not raised into the generic task-errored bucket.""" + for payload in ( + {"data": {"endpoints": None}}, + {"data": {"endpoints": "none"}}, + {"data": None}, + {}, + ): + provider = _FakeOpenRouterProvider( + models=[_model("m", canonical_slug="a/m")], db_id=2 + ) + _mock_transport( + monkeypatch, lambda request, p=payload: httpx.Response(200, json=p) + ) + # Must not raise. + await mp.refresh_model_paths([provider]) + + +@pytest.mark.asyncio +async def test_openrouter_shared_base_url_fetched_once( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Two providers on the same OpenRouter base URL share the per-cycle + endpoint cache instead of fetching byte-identical bodies twice.""" + native = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + generic = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=4, + ) + counter = _mock_transport( + monkeypatch, lambda request: _endpoints_response("Anthropic") + ) + + await mp.refresh_model_paths([native, generic]) + + assert counter["requests"] == 1 + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert "openrouter:Anthropic" in paths + assert "generic:Anthropic" in paths + + +@pytest.mark.asyncio +async def test_openrouter_fanout_is_bounded( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) + models = [_model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20)] + provider = _FakeOpenRouterProvider(models=models, db_id=2) + + state = {"current": 0, "max": 0} + + async def _slow_handler(request: httpx.Request) -> httpx.Response: + state["current"] += 1 + state["max"] = max(state["max"], state["current"]) + await asyncio.sleep(0.02) + state["current"] -= 1 + return _endpoints_response("X") + + def _factory() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(_slow_handler)) + + monkeypatch.setattr(mp, "_make_http_client", _factory) + + await mp.refresh_model_paths([provider]) + assert state["max"] > 0, "transport fake was never exercised" + assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" + + +# --------------------------------------------------------------------------- # +# Query endpoints +# --------------------------------------------------------------------------- # + + +async def _seed_two_provider_shared_model(engine: AsyncEngine) -> None: + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other/v1", + models=[_model("claude-opus-4.6")], + db_id=2, + ) + await mp.refresh_model_paths([p1, p2]) + + +@pytest.mark.asyncio +async def test_same_model_two_providers_two_paths( + patched_session: AsyncEngine, +) -> None: + await _seed_two_provider_shared_model(patched_session) + + payload = await mp.get_all_model_paths() + assert len(payload["data"]) == 1 + entry = payload["data"][0] + assert entry["id"] == "claude-opus-4.6" + assert {p["path"] for p in entry["paths"]} == {"anthropic", "generic"} + assert "canonical_id" not in entry + assert all("canonical_id" not in p for p in entry["paths"]) + + +@pytest.mark.asyncio +async def test_get_all_model_paths_deduplicates_visible_paths( + patched_session: AsyncEngine, +) -> None: + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("anthropic/claude-opus-4.6")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=2, + ) + await mp.refresh_model_paths([p1, p2]) + + assert (await mp.get_all_model_paths())["data"] == [ + {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + ] + + +@pytest.mark.asyncio +async def test_get_all_model_paths_is_deterministic( + patched_session: AsyncEngine, +) -> None: + """Output must not depend on rowid insertion order, which changes every + refresh cycle.""" + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("b-model"), _model("a-model")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + first = await mp.get_all_model_paths() + await mp.refresh_model_paths([provider]) + second = await mp.get_all_model_paths() + assert first["data"] == second["data"] + assert [e["id"] for e in first["data"]] == ["a-model", "b-model"] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_returns_only_paths( + patched_session: AsyncEngine, +) -> None: + await _seed_two_provider_shared_model(patched_session) + + payload = await mp.get_paths_for_model("claude-opus-4.6") + assert {p["path"] for p in payload["data"]} == {"anthropic", "generic"} + assert all(set(p.keys()) == {"path"} for p in payload["data"]) + assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("z-ai/glm-5v-turbo")], + db_id=4, + ) + await mp.refresh_model_paths([provider]) + + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ + {"path": "generic"} + ] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( + patched_session: AsyncEngine, +) -> None: + p1 = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("deepseek-v4-pro")], + db_id=7, + ) + p2 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("deepseek/deepseek-v4-pro")], + db_id=4, + ) + await mp.refresh_model_paths([p1, p2]) + + short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] + prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] + + assert {p["path"] for p in short_paths} == {"generic", "anthropic"} + assert prefixed_paths == short_paths + + +@pytest.mark.asyncio +async def test_get_paths_for_model_multi_segment_id_matches_models_listing( + patched_session: AsyncEngine, +) -> None: + """For three-segment ids the discovery id must be the same base id the + rest of the system exposes (first-slash rule), not the last segment.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("accounts/fireworks/models/glm-5")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} + assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ + {"path": "generic"} + ] + assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ + "data" + ] == [{"path": "generic"}] + + +# --------------------------------------------------------------------------- # +# Periodic refresh loop +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_refresh_loop_rereads_interval_and_picks_up_providers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 1, raising=False + ) + + seen_batches: list[list[Any]] = [] + + async def _fake_refresh(upstreams: list[Any]) -> None: + seen_batches.append(list(upstreams)) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + sleeps: list[float] = [] + + async def _fast_sleep(seconds: float) -> None: + sleeps.append(seconds) + if len(seen_batches) >= 2: + raise asyncio.CancelledError + + monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) + + batches = [["p1"], ["p1", "p2"]] + + def _provider() -> list[Any]: + return batches[min(len(seen_batches), len(batches) - 1)] + + await mp.refresh_model_paths_periodically(_provider) # type: ignore[arg-type] + + assert seen_batches[0] == ["p1"] + assert seen_batches[1] == ["p1", "p2"], "loop must re-resolve upstreams each cycle" + assert all(s >= 1 for s in sleeps) + + +@pytest.mark.asyncio +async def test_refresh_loop_idles_while_disabled( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A non-positive interval (or the kill switch) must idle the loop, not + exit it, so runtime re-enabling takes effect.""" + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + + refresh_calls: list[Any] = [] + + async def _fake_refresh(upstreams: list[Any]) -> None: + refresh_calls.append(upstreams) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + idle_sleeps: list[float] = [] + + async def _fast_sleep(seconds: float) -> None: + idle_sleeps.append(seconds) + if len(idle_sleeps) >= 2: + raise asyncio.CancelledError + + monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) + + await mp.refresh_model_paths_periodically(lambda: [object()]) + + assert refresh_calls == [], "disabled loop must not refresh" + assert len(idle_sleeps) == 2, "disabled loop must keep polling, not exit" + + +def test_refresh_interval_respects_kill_switch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + assert mp._refresh_interval_seconds() == 0 + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + assert mp._refresh_interval_seconds() == 600 + + +# --------------------------------------------------------------------------- # +# HTTP endpoints # --------------------------------------------------------------------------- # @@ -662,16 +1086,19 @@ def _make_model_paths_app() -> FastAPI: def test_model_paths_endpoint_returns_all_paths( monkeypatch: pytest.MonkeyPatch, ) -> None: - async def _fake_get_all_model_paths() -> list[dict[str, Any]]: - return [ - { - "id": "claude-opus-4.6", - "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"}, - ], - } - ] + async def _fake_get_all_model_paths() -> dict[str, Any]: + return { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + {"path": "anthropic"}, + {"path": "openrouter:Anthropic"}, + ], + } + ], + "updated_at": 1753500000, + } monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) @@ -687,7 +1114,8 @@ def test_model_paths_endpoint_returns_all_paths( {"path": "openrouter:Anthropic"}, ], } - ] + ], + "updated_at": 1753500000, } @@ -696,9 +1124,9 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) -> None: calls: list[str] = [] - async def _fake_get_paths_for_model(model_id: str) -> list[dict[str, Any]]: + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: calls.append(model_id) - return [{"path": "generic:Anthropic"}] + return {"data": [{"path": "generic:Anthropic"}], "updated_at": None} monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) @@ -708,5 +1136,8 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) assert response.status_code == 200 - assert response.json() == {"data": [{"path": "generic:Anthropic"}]} + assert response.json() == { + "data": [{"path": "generic:Anthropic"}], + "updated_at": None, + } assert calls == ["anthropic/claude-opus-4.6"] From 06dba681c53250e038aa8434517df00eef6c7802 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 13:26:49 +0200 Subject: [PATCH 31/46] fix: satisfy strict mypy in model-paths tests Replace untyped lambdas with typed handler/provider functions; CI runs mypy over tests as well. --- tests/unit/test_model_paths.py | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index b1747436..7d3393fe 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -13,7 +13,7 @@ import json import os from contextlib import asynccontextmanager from types import SimpleNamespace -from typing import Any, AsyncGenerator, Callable +from typing import Any, AsyncGenerator, Callable, cast import httpx import pytest @@ -760,9 +760,13 @@ async def test_openrouter_bad_payload_shapes_do_not_raise( provider = _FakeOpenRouterProvider( models=[_model("m", canonical_slug="a/m")], db_id=2 ) - _mock_transport( - monkeypatch, lambda request, p=payload: httpx.Response(200, json=p) - ) + + def _handler( + request: httpx.Request, p: dict[str, Any] | None = payload + ) -> httpx.Response: + return httpx.Response(200, json=p) + + _mock_transport(monkeypatch, _handler) # Must not raise. await mp.refresh_model_paths([provider]) @@ -1051,7 +1055,10 @@ async def test_refresh_loop_idles_while_disabled( monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) - await mp.refresh_model_paths_periodically(lambda: [object()]) + def _upstreams() -> list[BaseUpstreamProvider]: + return [cast(BaseUpstreamProvider, object())] + + await mp.refresh_model_paths_periodically(_upstreams) assert refresh_calls == [], "disabled loop must not refresh" assert len(idle_sleeps) == 2, "disabled loop must keep polling, not exit" From c5da73f1e97bbd670c48ea48a5c81fee27b1677e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 20:12:12 +0200 Subject: [PATCH 32/46] revert: remove unrelated repository formatting --- routstr/algorithm.py | 20 +- routstr/balance.py | 29 +-- routstr/core/admin.py | 26 +- routstr/core/log_manager.py | 21 +- routstr/core/usage_analytics_store.py | 29 +-- routstr/nostr/analytics.py | 8 +- routstr/payment/cost_calculation.py | 19 +- routstr/payment/usage.py | 4 +- routstr/upstream/azure.py | 4 +- routstr/upstream/ehbp.py | 32 ++- routstr/upstream/gemini.py | 4 +- routstr/upstream/gemini_messages.py | 4 +- routstr/upstream/groq.py | 4 +- routstr/upstream/litellm_routing.py | 4 +- routstr/upstream/messages_dispatch.py | 12 +- routstr/upstream/model_paths.py | 322 +++++++++++++++---------- routstr/upstream/ollama.py | 8 +- routstr/upstream/rate_limit.py | 4 +- routstr/upstream/request_correction.py | 4 +- routstr/upstream/routstr.py | 3 +- routstr/upstream/xai.py | 4 +- routstr/wallet.py | 13 +- 22 files changed, 319 insertions(+), 259 deletions(-) diff --git a/routstr/algorithm.py b/routstr/algorithm.py index ef4b8574..fbc5388e 100644 --- a/routstr/algorithm.py +++ b/routstr/algorithm.py @@ -232,10 +232,7 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if ( - model_to_use.forwarded_model_id - and model_to_use.forwarded_model_id not in aliases - ): + if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: aliases.append(model_to_use.forwarded_model_id) # Try to set each alias @@ -325,10 +322,7 @@ def create_model_mappings( aliases.append(prefixed_id) # Register forwarded_model_id as a routable alias - if ( - model_to_use.forwarded_model_id - and model_to_use.forwarded_model_id not in aliases - ): + if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases: aliases.append(model_to_use.forwarded_model_id) for alias in aliases: @@ -348,10 +342,16 @@ def create_model_mappings( forwarded_model_ids, the one whose forwarded_model_id equals the requested alias wins. """ - if model.forwarded_model_id and model.forwarded_model_id.lower() == alias: + if ( + model.forwarded_model_id + and model.forwarded_model_id.lower() == alias + ): return 5 - if model.id and model.id.lower() == alias: + if ( + model.id + and model.id.lower() == alias + ): return 4 model_base = get_base_model_id(model.id) diff --git a/routstr/balance.py b/routstr/balance.py index 03dc4d33..91b19ce5 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -260,11 +260,7 @@ async def _lookup_key_no_create( async def _restore_balance( - session: AsyncSession, - hashed_key: str, - balance: int, - reserved_balance: int, - mint_url: str, + session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str ) -> None: """Restore balance after a failed refund mint attempt.""" restore_stmt = ( @@ -279,11 +275,7 @@ async def _restore_balance( await session.commit() logger.info( "refund_wallet_endpoint: balance restored after mint failure", - extra={ - "hashed_key": hashed_key, - "restored_balance": balance, - "mint_url": mint_url, - }, + extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url}, ) @@ -468,23 +460,11 @@ async def refund_wallet_endpoint( except HTTPException: # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) + await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") raise except Exception as e: # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) + await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "") error_msg = str(e) logger.error( "refund_wallet_endpoint: mint/send failed", @@ -705,6 +685,7 @@ async def reset_child_key_spent( return {"success": True, "message": "Child key balance reset successfully."} + @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 1521510e..66a1d288 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -68,9 +68,7 @@ async def require_admin_api(request: Request) -> None: async with create_session() as session: result = await session.exec(select(CliToken).where(CliToken.token == token)) cli_token = result.first() - if cli_token and ( - cli_token.expires_at is None or cli_token.expires_at > now_ts - ): + if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts): cli_token.last_used_at = now_ts session.add(cli_token) await session.commit() @@ -257,12 +255,16 @@ async def update_password(request: Request, password_update: PasswordUpdate) -> secret = await get_secret(session) if not secret.admin_password_hash: - raise HTTPException(status_code=500, detail="Admin password not configured") + raise HTTPException( + status_code=500, detail="Admin password not configured" + ) if not vault.verify_password( password_update.current_password, secret.admin_password_hash ): - raise HTTPException(status_code=401, detail="Current password is incorrect") + raise HTTPException( + status_code=401, detail="Current password is incorrect" + ) # Validate new password new_password = password_update.new_password.strip() @@ -978,7 +980,9 @@ async def update_upstream_provider_by_slug( lookup = _validate_slug(payload.slug) async with create_session() as session: result = await session.exec( - select(UpstreamProviderRow).where(UpstreamProviderRow.slug == lookup) + select(UpstreamProviderRow).where( + UpstreamProviderRow.slug == lookup + ) ) provider = result.first() if not provider: @@ -1665,11 +1669,7 @@ async def get_transactions_api( ) total = count_result.one() - stmt = ( - base.order_by(col(CashuTransaction.created_at).desc()) - .offset(offset) - .limit(limit) - ) + stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit) results = await session.exec(stmt) transactions = results.all() @@ -1679,7 +1679,9 @@ async def get_transactions_api( } -@admin_router.get("/api/lightning-invoices", dependencies=[Depends(require_admin_api)]) +@admin_router.get( + "/api/lightning-invoices", dependencies=[Depends(require_admin_api)] +) async def get_lightning_invoices_api( status: str | None = None, purpose: str | None = None, diff --git a/routstr/core/log_manager.py b/routstr/core/log_manager.py index b111f68a..0444dcbf 100644 --- a/routstr/core/log_manager.py +++ b/routstr/core/log_manager.py @@ -408,9 +408,7 @@ class LogManager: def get_error_details(self, hours: int = 24, limit: int = 100) -> dict: def compute() -> dict: try: - return self._usage_store.get_error_details( - hours_back=hours, limit=limit - ) + return self._usage_store.get_error_details(hours_back=hours, limit=limit) except Exception as e: logger.error( f"Usage analytics index failed, falling back to log scan: {e}" @@ -630,7 +628,8 @@ class LogManager: stats["total_tokens"] += input_tokens + output_tokens failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: stats["total_requests"] += 1 @@ -788,9 +787,7 @@ class LogManager: if bucket_key: model_mix_buckets[bucket_key][model] += 1 if revenue_msats > 0: - model_mix_revenue_buckets[bucket_key][model] += ( - revenue_msats - ) + model_mix_revenue_buckets[bucket_key][model] += revenue_msats model_mix_revenue_totals[model] += revenue_msats if input_tokens > 0 or output_tokens > 0: token_total = input_tokens + output_tokens @@ -804,7 +801,8 @@ class LogManager: bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: summary_stats["total_requests"] += 1 @@ -874,7 +872,9 @@ class LogManager: models.sort(key=lambda x: float(x["net_revenue_sats"]), reverse=True) latest_errors = [ item - for _, item in sorted(latest_errors_heap, key=lambda x: x[0], reverse=True) + for _, item in sorted( + latest_errors_heap, key=lambda x: x[0], reverse=True + ) ] top_model_limit = max(1, min(model_limit, 20)) top_models_requests = [ @@ -1051,7 +1051,8 @@ class LogManager: bucket["warnings"] += 1 failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: bucket["total_requests"] += 1 diff --git a/routstr/core/usage_analytics_store.py b/routstr/core/usage_analytics_store.py index 36fa4bcb..7ba90e24 100644 --- a/routstr/core/usage_analytics_store.py +++ b/routstr/core/usage_analytics_store.py @@ -314,7 +314,9 @@ class UsageAnalyticsStore: if column in existing_columns: return - conn.execute(f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}") + conn.execute( + f"ALTER TABLE {table} ADD COLUMN {column} {column_definition}" + ) logger.info(f"Migrated analytics schema: added {table}.{column}") def _drop_index_tables_locked(self, conn: sqlite3.Connection) -> None: @@ -362,11 +364,7 @@ class UsageAnalyticsStore: self._drop_index_tables_locked(conn) self._initialize_schema_locked(conn) - files = ( - log_files - if log_files is not None - else sorted(self.logs_dir.glob("app_*.log")) - ) + files = log_files if log_files is not None else sorted(self.logs_dir.glob("app_*.log")) for log_file in files: try: self._process_log_file_locked(conn, log_file, force_full_read=True) @@ -570,7 +568,8 @@ class UsageAnalyticsStore: model_bucket["revenue_msats"] += revenue_msats failed = ( - "upstream request failed" in message or "revert payment" in message + "upstream request failed" in message + or "revert payment" in message ) if failed: bucket["total_requests"] += 1 @@ -593,9 +592,9 @@ class UsageAnalyticsStore: if isinstance(max_cost, (int, float)) and max_cost > 0: max_cost_float = float(max_cost) bucket["refunds_msats"] += max_cost_float - model_updates[(minute_key, model)]["refunds_msats"] += ( - max_cost_float - ) + model_updates[(minute_key, model)][ + "refunds_msats" + ] += max_cost_float return ( end_offset, @@ -1033,9 +1032,7 @@ class UsageAnalyticsStore: """, (cutoff_timestamp,), ).fetchone() - total_error_count = ( - int(total_error_count_row[0]) if total_error_count_row else 0 - ) + total_error_count = int(total_error_count_row[0]) if total_error_count_row else 0 return { "errors": [ @@ -1207,7 +1204,11 @@ class UsageAnalyticsStore: total_successful = int(row["total_successful"]) total_revenue_msats = float(row["total_revenue_msats"]) total_tokens = int(row["total_tokens"]) - if total_successful <= 0 and total_revenue_msats <= 0 and total_tokens <= 0: + if ( + total_successful <= 0 + and total_revenue_msats <= 0 + and total_tokens <= 0 + ): continue bucket_ts = str(row["bucket_ts"]) diff --git a/routstr/nostr/analytics.py b/routstr/nostr/analytics.py index 8b6da590..e568b5e0 100644 --- a/routstr/nostr/analytics.py +++ b/routstr/nostr/analytics.py @@ -215,9 +215,7 @@ def _build_window_payload( summary = dashboard.get("summary", {}) model_usage_mix = dashboard.get("model_usage_mix", {}) - summary_payload = _build_summary_payload( - summary if isinstance(summary, dict) else {} - ) + summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {}) usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {} top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload) @@ -340,9 +338,7 @@ async def publish_usage_analytics() -> None: nsec = (settings.nsec or "").strip() if not nsec: if not warned_missing_nsec: - logger.info( - "NSEC is not configured; skipping analytics sharing to Nostr" - ) + logger.info("NSEC is not configured; skipping analytics sharing to Nostr") warned_missing_nsec = True await asyncio.sleep(DISABLED_POLL_SECONDS) continue diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index e7cee8ca..37ac15d3 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -224,7 +224,9 @@ async def calculate_cost( "Token counts %s in the upstream response but cannot be " "priced; the request will appear in dashboards with the " "raw counts and a fixed max-cost charge.", - "are present" if (input_tokens > 0 or output_tokens > 0) else "are zero", + "are present" + if (input_tokens > 0 or output_tokens > 0) + else "are zero", extra={ "base_cost_msats": max_cost, "model": response_data.get("model", "unknown"), @@ -301,7 +303,9 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float: # actually deducts from the balance. For non-BYOK providers (e.g. # OpenRouter) usage.cost already equals upstream_inference_cost, so we # fall through to the normal ``cost`` lookup below. - upstream_cost = _coerce_usd(cost_details.get("upstream_inference_cost")) + upstream_cost = _coerce_usd( + cost_details.get("upstream_inference_cost") + ) if upstream_cost > 0 and usage_data.get("is_byok"): byok_fee = _coerce_usd(usage_data.get("cost")) return upstream_cost + byok_fee @@ -332,7 +336,8 @@ def _get_pricing_rates( ``None`` means configured fixed pricing should be used by the caller. """ if settings.fixed_pricing and ( - settings.fixed_per_1k_input_tokens or settings.fixed_per_1k_output_tokens + settings.fixed_per_1k_input_tokens + or settings.fixed_per_1k_output_tokens ): return None @@ -388,8 +393,12 @@ def _get_pricing_rates( usd_per_sat = sats_usd_price() mspp_1k = input_usd * provider_fee * 1_000_000.0 / usd_per_sat mspc_1k = output_usd * provider_fee * 1_000_000.0 / usd_per_sat - cache_read_usd = _coerce_usd(pricing.get("cache_read_input_token_cost")) - cache_write_usd = _coerce_usd(pricing.get("cache_creation_input_token_cost")) + cache_read_usd = _coerce_usd( + pricing.get("cache_read_input_token_cost") + ) + cache_write_usd = _coerce_usd( + pricing.get("cache_creation_input_token_cost") + ) mscr_1k = ( cache_read_usd * provider_fee * 1_000_000.0 / usd_per_sat if cache_read_usd > 0 diff --git a/routstr/payment/usage.py b/routstr/payment/usage.py index 11d5c01e..02c90055 100644 --- a/routstr/payment/usage.py +++ b/routstr/payment/usage.py @@ -110,7 +110,9 @@ def normalize_usage(usage_data: object) -> NormalizedUsage | None: if not isinstance(usage_data, dict): return None - output_tokens = _first_token_count(usage_data, "completion_tokens", "output_tokens") + output_tokens = _first_token_count( + usage_data, "completion_tokens", "output_tokens" + ) cache_read, cache_write = _extract_cache_tokens(usage_data) # ``prompt_tokens`` is the inclusive grand total; ``input_tokens`` (Anthropic diff --git a/routstr/upstream/azure.py b/routstr/upstream/azure.py index 985bcfd2..a693b763 100644 --- a/routstr/upstream/azure.py +++ b/routstr/upstream/azure.py @@ -94,7 +94,9 @@ class AzureUpstreamProvider(BaseUpstreamProvider): deployment_id = deployment_id.split("/")[-1] return f"openai/deployments/{deployment_id}/{clean_path}" - def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: + def get_request_base_url( + self, path: str, model_obj: "Model | None" = None + ) -> str: """Use endpoint root, stripping accidental /openai/v1 suffix if present.""" base_url = self.base_url.rstrip("/") marker = "/openai/v1" diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index a3d1505a..96955492 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -191,9 +191,7 @@ def _resolve_ehbp_target_url( otherwise the header is ignored so callers cannot redirect other providers or leak upstream API keys. """ - override_header = ( - profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER - ) + override_header = profile.client_target_url_header if profile else _ENCLAVE_URL_HEADER if not override_header: return target_url enclave_url = _get_header_case_insensitive(headers, override_header) @@ -297,7 +295,9 @@ def _build_cost_info( return result -def _inject_cost_response_headers(headers: dict[str, str], cost_info: dict) -> None: +def _inject_cost_response_headers( + headers: dict[str, str], cost_info: dict +) -> None: """Add per-request cost headers to an EHBP response. Since EHBP response bodies are opaque encrypted blobs, cost cannot be @@ -375,7 +375,9 @@ async def _compute_ehbp_actual_cost( resolved_upstream_model = ( actual_model_obj.forwarded_model_id or actual_model_obj.id ) - resolved_identity = _normalize_upstream_model_id(resolved_upstream_model) + resolved_identity = _normalize_upstream_model_id( + resolved_upstream_model + ) if resolved_identity != expected_identity: logger.info( "EHBP served model differs from requested, using actual " @@ -515,9 +517,7 @@ async def finalize_ehbp_actual_cost_payment( billing_key = await get_billing_key(key, session) key_hash = key.hashed_key billing_key_hash = billing_key.hashed_key - total_cost_msats = max( - 0, int(cost_info.get("total_msats", reserved_cost_for_model)) - ) + total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model))) now = int(time.time()) safe_reserved = case( @@ -560,9 +560,7 @@ async def finalize_ehbp_actual_cost_payment( ) child_result = await session.exec(child_stmt) # type: ignore[call-overload] - if result.rowcount == 0 or ( - child_result is not None and child_result.rowcount == 0 - ): + if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): await session.rollback() logger.error( "Failed to finalize EHBP usage-based payment", @@ -692,9 +690,7 @@ async def finalize_ehbp_max_cost_payment( else: child_result = None - if result.rowcount == 0 or ( - child_result is not None and child_result.rowcount == 0 - ): + if result.rowcount == 0 or (child_result is not None and child_result.rowcount == 0): await session.rollback() logger.error( "Failed to finalize EHBP max-cost payment", @@ -1038,9 +1034,7 @@ async def forward_ehbp_x_cashu_request( target_url = _resolve_ehbp_target_url( target.url, path, headers, provider_type, profile ) - upstream_headers = _prepare_ehbp_upstream_headers( - headers, target.headers, profile - ) + upstream_headers = _prepare_ehbp_upstream_headers(headers, target.headers, profile) request_body = await request.body() # Merge query params into the target URL @@ -1088,7 +1082,9 @@ async def forward_ehbp_x_cashu_request( usage_source = ( "header" if usage_header_name - and any(k.lower() == usage_header_name.lower() for k, _ in resp.headers) + and any( + k.lower() == usage_header_name.lower() for k, _ in resp.headers + ) else ("trailer" if usage_header else "none") ) diff --git a/routstr/upstream/gemini.py b/routstr/upstream/gemini.py index d58199c2..54de41a3 100644 --- a/routstr/upstream/gemini.py +++ b/routstr/upstream/gemini.py @@ -94,7 +94,9 @@ class GeminiUpstreamProvider(BaseUpstreamProvider): """ return self.base_url.rstrip("/").removesuffix("/openai") + "/openai" - def get_request_base_url(self, path: str, model_obj: "Model | None" = None) -> str: + def get_request_base_url( + self, path: str, model_obj: "Model | None" = None + ) -> str: """Route every proxied request to the OpenAI-compat surface. Required because the stored ``base_url`` typically points at the diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index a7e41c9b..11440b91 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -371,7 +371,9 @@ async def dispatch_gemini_messages( aggregates). """ if not request_body: - raise UpstreamError("Missing request body for /v1/messages", status_code=400) + raise UpstreamError( + "Missing request body for /v1/messages", status_code=400 + ) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/groq.py b/routstr/upstream/groq.py index a0c9475e..17103c35 100644 --- a/routstr/upstream/groq.py +++ b/routstr/upstream/groq.py @@ -20,9 +20,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row( - cls, provider_row: "UpstreamProviderRow" - ) -> "GroqUpstreamProvider": + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "GroqUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/upstream/litellm_routing.py b/routstr/upstream/litellm_routing.py index b7790394..0b2a92a3 100644 --- a/routstr/upstream/litellm_routing.py +++ b/routstr/upstream/litellm_routing.py @@ -91,7 +91,9 @@ OLLAMA_HOST_HINTS: tuple[str, ...] = ( ) -def detect_litellm_prefix(base_url: str | None, default: str = DEFAULT_PREFIX) -> str: +def detect_litellm_prefix( + base_url: str | None, default: str = DEFAULT_PREFIX +) -> str: """Return the litellm provider prefix (`"/"`) for `base_url`. Falls back to `default` when the host doesn't match any known provider. diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index efcf591b..3d689922 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -108,7 +108,9 @@ def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]: return events, buffer -def events_from_chunk(chunk: object, sse_buffer: bytes) -> tuple[list[dict], bytes]: +def events_from_chunk( + chunk: object, sse_buffer: bytes +) -> tuple[list[dict], bytes]: """Normalize a stream chunk into one or more event dicts. ``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE @@ -199,7 +201,9 @@ async def aggregate_anthropic_events_to_message( raw_json = partial_json.pop(idx, None) if raw_json is not None and idx < len(blocks): try: - blocks[idx]["input"] = json.loads(raw_json) if raw_json else {} + blocks[idx]["input"] = ( + json.loads(raw_json) if raw_json else {} + ) except json.JSONDecodeError: blocks[idx]["input"] = raw_json elif etype == "message_delta": @@ -441,7 +445,9 @@ async def dispatch_anthropic_messages( on bad input or upstream failure. """ if not request_body: - raise UpstreamError("Missing request body for /v1/messages", status_code=400) + raise UpstreamError( + "Missing request body for /v1/messages", status_code=400 + ) try: body: dict = json.loads(request_body) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index beaafa6a..f3bb9971 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -1,19 +1,15 @@ """Model-path discovery service. -Exposes every upstream provider path a Routstr model is reachable through. -This is discovery/visibility data only — routing still selects the cheapest or -best provider separately. +Exposes every selectable upstream route a Routstr model is reachable through. +This PR remains discovery-only: request-side routing will consume the opaque +selectors in a follow-up. -A *path* is the provider string that may appear in Routstr chat completion -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: +A path is a standard percent-encoded query string containing the normalized +upstream URL and, for an exact OpenRouter endpoint, its machine-readable tag. +Display names never participate in identity:: -- Direct upstream -> ```` e.g. ``anthropic`` -- Generic/custom OpenRouter-compatible upstream -> ``generic:`` -- Native OpenRouter routing to a sub-provider -> ``openrouter:`` -- Native OpenRouter with no usable sub-provider -> ``unknown`` + url=https%3A%2F%2Fapi.anthropic.com%2Fv1 + url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider=google-vertex%2Fus """ from __future__ import annotations @@ -21,10 +17,12 @@ from __future__ import annotations import asyncio import random import time -from typing import TYPE_CHECKING, Callable +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable +from urllib.parse import urlencode import httpx -from sqlalchemy import insert, or_ +from sqlalchemy import insert from sqlalchemy.orm import selectinload from sqlmodel import col, delete, select @@ -51,6 +49,46 @@ _PERSIST_CHUNK_SIZE = 500 ModelKey = tuple[str, int] +@dataclass(frozen=True) +class EndpointIdentity: + """Exact OpenRouter endpoint identity returned by ``/endpoints``.""" + + tag: str + provider_name: str | None + + +@dataclass(frozen=True) +class DiscoveredPath: + """One model route ready for persistence and API serialization.""" + + model_id: str + path: str + upstream_url: str + provider_tag: str | None = None + provider_name: str | None = None + + +@dataclass(frozen=True) +class ProviderPathSnapshot: + """Refresh result plus model IDs whose prior rows must survive degradation.""" + + paths: tuple[DiscoveredPath, ...] + preserve_model_ids: frozenset[str] = frozenset() + + +def normalize_upstream_url(base_url: str) -> str: + """Normalize route identity without changing URL semantics.""" + return base_url.rstrip("/") + + +def encode_model_path(base_url: str, provider_tag: str | None = None) -> str: + """Encode a stable opaque selector for future request-side routing.""" + components = [("url", normalize_upstream_url(base_url))] + if provider_tag: + components.append(("provider", provider_tag)) + return urlencode(components) + + def _make_http_client() -> httpx.AsyncClient: """Client factory, separated so tests can substitute a mock transport.""" return httpx.AsyncClient() @@ -69,9 +107,15 @@ def is_openrouter_base_url(base_url: str | None) -> bool: def exposed_model_id(model: object) -> str: - """Client-visible ``/v1/models`` id for a cached model.""" + """Return exactly the ID advertised by ``/v1/models``. + + A forwarded ID is already a public routable alias and must remain intact, + including any slash. Without one, ``/v1/models`` exposes the base ID. + """ forwarded = getattr(model, "forwarded_model_id", None) - return forwarded or getattr(model, "id") + if forwarded: + return forwarded + return public_model_id(getattr(model, "id")) def public_model_id(model_id: str) -> str: @@ -116,7 +160,7 @@ class _RefreshCycleState: """ def __init__(self) -> None: - self.endpoint_cache: dict[tuple[str, str], list[str] | None] = {} + self.endpoint_cache: dict[tuple[str, str], list[EndpointIdentity] | None] = {} self.rate_limited = False @@ -127,8 +171,8 @@ async def _fetch_openrouter_endpoint_subproviders( author_slug: str, semaphore: asyncio.Semaphore, cycle: _RefreshCycleState, -) -> list[str] | None: - """Return sub-provider names for one model, or ``None`` when unknown. +) -> list[EndpointIdentity] | None: + """Return exact endpoint identities for one model, or ``None`` when unknown. ``None`` (not ``[]``) signals a degraded fetch — network failure, rate limit, non-200, or an unparseable payload — so callers can distinguish @@ -143,7 +187,7 @@ async def _fetch_openrouter_endpoint_subproviders( url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} - result: list[str] | None + result: list[EndpointIdentity] | None async with semaphore: try: resp = await client.get( @@ -177,14 +221,24 @@ async def _fetch_openrouter_endpoint_subproviders( endpoints = resp.json().get("data", {}).get("endpoints", []) if not isinstance(endpoints, list): endpoints = [] - names: list[str] = [] + identities: dict[str, EndpointIdentity] = {} for endpoint in endpoints: - provider_name = ( - endpoint.get("provider_name") if isinstance(endpoint, dict) else None + if not isinstance(endpoint, dict): + continue + tag = endpoint.get("tag") + if not isinstance(tag, str) or not tag.strip(): + continue + provider_name = endpoint.get("provider_name") + identities.setdefault( + tag, + EndpointIdentity( + tag=tag, + provider_name=provider_name + if isinstance(provider_name, str) and provider_name + else None, + ), ) - if provider_name: - names.append(provider_name) - result = list(dict.fromkeys(names)) + result = list(identities.values()) except Exception as e: # noqa: BLE001 logger.warning( "OpenRouter endpoint discovery bad payload", @@ -287,46 +341,41 @@ async def _collect_provider_paths( 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. +) -> ProviderPathSnapshot: + """Collect selectable routes while marking model-level degraded fetches. - 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. + A failed OpenRouter lookup preserves only that model's prior rows. Other + models in the same provider still refresh, so a partial outage cannot erase + valid discovery data or freeze the entire provider snapshot. """ cycle = cycle or _RefreshCycleState() models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) - base_paths = upstream.discovery_base_paths() + upstream_url = normalize_upstream_url(upstream.base_url) + + def _base_path(model: object) -> DiscoveredPath: + return DiscoveredPath( + model_id=exposed_model_id(model), + path=encode_model_path(upstream_url), + upstream_url=upstream_url, + ) if not is_openrouter_base_url(upstream.base_url): - return [ - (exposed_model_id(model), path) for model in models for path in base_paths - ] + return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models)) if not (upstream.provider_type or "").strip(): - return [] + return ProviderPathSnapshot(paths=()) - any_fetch_succeeded = False - any_fetch_attempted = False semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) async with _make_http_client() as client: - async def _for_model(model: object) -> list[tuple[str, str]]: - nonlocal any_fetch_succeeded, any_fetch_attempted + async def _for_model( + model: object, + ) -> tuple[list[DiscoveredPath], str | None]: 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 pairs - any_fetch_attempted = True - sub_providers = await _fetch_openrouter_endpoint_subproviders( + return [_base_path(model)], None + endpoints = await _fetch_openrouter_endpoint_subproviders( client, upstream.base_url, upstream.api_key, @@ -334,67 +383,78 @@ async def _collect_provider_paths( semaphore, cycle, ) - 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)) + if endpoints is None: + return [], model_id + paths = [_base_path(model)] + paths.extend( + DiscoveredPath( + model_id=model_id, + path=encode_model_path(upstream_url, endpoint.tag), + upstream_url=upstream_url, + provider_tag=endpoint.tag, + provider_name=endpoint.provider_name, + ) + for endpoint in endpoints + ) + return paths, None results = await asyncio.gather( - *(_for_model(m) for m in models), return_exceptions=True + *(_for_model(model) for model 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: + paths: list[DiscoveredPath] = [] + preserve_model_ids: set[str] = set() + for model, result in zip(models, results): if isinstance(result, BaseException): + model_id = exposed_model_id(model) + preserve_model_ids.add(model_id) logger.warning( "OpenRouter endpoint discovery task errored", extra={"provider": upstream.provider_type, "error": str(result)}, ) continue - pairs.extend(result) + model_paths, preserved_model_id = result + paths.extend(model_paths) + if preserved_model_id: + preserve_model_ids.add(preserved_model_id) - return pairs + return ProviderPathSnapshot( + paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids) + ) async def _persist_provider_paths( - upstream_provider_id: int, pairs: list[tuple[str, str]] + upstream_provider_id: int, snapshot: ProviderPathSnapshot ) -> None: - """Replace all rows for ``upstream_provider_id`` with ``pairs``. - - Replacement (not upsert) so stale paths disappear when provider config or - 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)) + """Replace refreshed rows while retaining model-level degraded snapshots.""" + unique_paths = list( + {(path.model_id, path.path): path for path in snapshot.paths}.values() + ) now = int(time.time()) async with create_session() as session: - await session.exec( # type: ignore[call-overload] - delete(ModelPathRow).where( - col(ModelPathRow.upstream_provider_id) == upstream_provider_id - ) + delete_stmt = delete(ModelPathRow).where( + col(ModelPathRow.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] + if snapshot.preserve_model_ids: + delete_stmt = delete_stmt.where( + col(ModelPathRow.model_id).not_in(sorted(snapshot.preserve_model_ids)) + ) + await session.exec(delete_stmt) # type: ignore[call-overload] + for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE): + chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE] await session.execute( insert(ModelPathRow), [ { - "model_id": model_id, - "path": path, + "model_id": discovered.model_id, + "path": discovered.path, + "upstream_url": discovered.upstream_url, + "provider_tag": discovered.provider_tag, + "provider_name": discovered.provider_name, "upstream_provider_id": upstream_provider_id, "updated_at": now, } - for model_id, path in chunk + for discovered in chunk ], ) await session.commit() @@ -454,22 +514,22 @@ async def refresh_model_paths( if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: continue try: - pairs = await _collect_provider_paths( + snapshot = await _collect_provider_paths( upstream, overrides_by_key=overrides_by_key, disabled_model_keys=disabled_model_keys, cycle=cycle, ) - if pairs is None: + if snapshot.preserve_model_ids: logger.warning( - "Model paths unknown this cycle; keeping previous rows", + "Some model paths are unknown; keeping their previous rows", extra={ "provider": upstream.provider_type or upstream.base_url, "db_id": upstream.db_id, + "preserved_models": len(snapshot.preserve_model_ids), }, ) - continue - await _persist_provider_paths(upstream.db_id, pairs) + await _persist_provider_paths(upstream.db_id, snapshot) except Exception as e: # noqa: BLE001 - isolate per-provider failures logger.error( "Failed to refresh model paths for provider", @@ -482,6 +542,21 @@ 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.""" + from ..proxy import get_upstreams + + matching = [ + upstream + for upstream in get_upstreams() + if upstream.db_id == upstream_provider_id + ] + if matching: + await refresh_model_paths(matching) + else: + await prune_model_paths_for_inactive_providers() + + def _refresh_interval_seconds() -> int: """Current interval, re-read every loop so runtime setting changes apply.""" from ..core.settings import settings @@ -535,8 +610,19 @@ async def refresh_model_paths_periodically( break +def _serialize_path(row: ModelPathRow) -> dict[str, Any]: + provider = None + if row.provider_tag or row.provider_name: + provider = {"name": row.provider_name, "slug": row.provider_tag} + return { + "path": row.path, + "upstream_url": row.upstream_url, + "provider": provider, + } + + async def get_all_model_paths() -> dict: - """All models with their paths, shaped for ``GET /v1/models/paths``.""" + """All models with their exact selectable routes.""" async with create_session() as session: rows = ( await session.exec( @@ -548,49 +634,37 @@ async def get_all_model_paths() -> dict: ) ).all() - grouped: dict[str, list[dict]] = {} + grouped: dict[str, list[dict[str, Any]]] = {} seen_paths: dict[str, set[str]] = {} updated_at = 0 for row in rows: updated_at = max(updated_at, row.updated_at) - model_id = public_model_id(row.model_id) - if row.path in seen_paths.setdefault(model_id, set()): + if row.path in seen_paths.setdefault(row.model_id, set()): continue - seen_paths[model_id].add(row.path) - grouped.setdefault(model_id, []).append({"path": row.path}) - # 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}) + seen_paths[row.model_id].add(row.path) + grouped.setdefault(row.model_id, []).append(_serialize_path(row)) + data = [ + { + "id": grouped_model_id, + "paths": sorted( + grouped[grouped_model_id], key=lambda item: str(item["path"]) + ), + } + for grouped_model_id in sorted(grouped) + ] return {"data": data, "updated_at": updated_at or None} async def get_paths_for_model(model_id: str) -> dict: - """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. The candidate set is - narrowed in SQL (exact id or ``%/`` suffix) so the route does not - materialize the whole table per request. - """ - # 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)} + """Return paths only for the exact model ID advertised by ``/v1/models``.""" 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) - .where(or_(*conditions)) + .where(col(ModelPathRow.model_id) == model_id) .order_by( col(ModelPathRow.path), col(ModelPathRow.upstream_provider_id), - col(ModelPathRow.model_id), ) ) ).all() @@ -599,15 +673,9 @@ async def get_paths_for_model(model_id: str) -> dict: paths: list[dict] = [] updated_at = 0 for row in rows: - # 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}) + paths.append(_serialize_path(row)) return {"data": paths, "updated_at": updated_at or None} diff --git a/routstr/upstream/ollama.py b/routstr/upstream/ollama.py index c4873ea0..9fed0154 100644 --- a/routstr/upstream/ollama.py +++ b/routstr/upstream/ollama.py @@ -66,7 +66,9 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): """Strip 'ollama/' prefix for Ollama API compatibility.""" return model_id.removeprefix("ollama/") - def get_request_base_url(self, path: str, model_obj: Model | None = None) -> str: + def get_request_base_url( + self, path: str, model_obj: Model | None = None + ) -> str: """Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint.""" return f"{self.base_url.rstrip('/')}/v1" @@ -183,9 +185,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider): except Exception: self._models_cache = models_with_fees - self._models_by_id = { - m.forwarded_model_id or m.id: m for m in self._models_cache - } + self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache} logger.info( f"Refreshed models cache for {self.base_url}", extra={"model_count": len(models)}, diff --git a/routstr/upstream/rate_limit.py b/routstr/upstream/rate_limit.py index ac1eff78..dca1ba5b 100644 --- a/routstr/upstream/rate_limit.py +++ b/routstr/upstream/rate_limit.py @@ -119,9 +119,7 @@ def classify_rate_limit( retry_match = _RETRY_RE.search(redacted) if retry_match is not None: value = float(retry_match.group(1)) - retry_after = ( - value / 1000.0 if retry_match.group(2).lower() == "ms" else value - ) + retry_after = value / 1000.0 if retry_match.group(2).lower() == "ms" else value limit_name_match = _LIMIT_NAME_RE.search(redacted) diff --git a/routstr/upstream/request_correction.py b/routstr/upstream/request_correction.py index 8e4379a0..c2ea5b1d 100644 --- a/routstr/upstream/request_correction.py +++ b/routstr/upstream/request_correction.py @@ -84,7 +84,9 @@ def extract_error_message(response: Response) -> str: return "" -def strip_unsupported_param(body: dict, error_message: str) -> tuple[dict, str] | None: +def strip_unsupported_param( + body: dict, error_message: str +) -> tuple[dict, str] | None: """Drop a top-level param the upstream named as unsupported/deprecated. Returns ``(new_body, param)`` (a new dict, original untouched) when the diff --git a/routstr/upstream/routstr.py b/routstr/upstream/routstr.py index de1aa3bd..0371946a 100644 --- a/routstr/upstream/routstr.py +++ b/routstr/upstream/routstr.py @@ -50,7 +50,8 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider): def normalize_request_path( self, path: str, model_obj: "Model | None" = None ) -> str: - """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.""" + """Preserve the ``v1/`` prefix when forwarding to an upstream Routstr. + """ return path.lstrip("/") @classmethod diff --git a/routstr/upstream/xai.py b/routstr/upstream/xai.py index 12e3dd93..58caaba0 100644 --- a/routstr/upstream/xai.py +++ b/routstr/upstream/xai.py @@ -21,9 +21,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider): ) @classmethod - def _build_from_row( - cls, provider_row: "UpstreamProviderRow" - ) -> "XAIUpstreamProvider": + def _build_from_row(cls, provider_row: "UpstreamProviderRow") -> "XAIUpstreamProvider": return cls( api_key=provider_row.api_key, provider_fee=provider_row.provider_fee, diff --git a/routstr/wallet.py b/routstr/wallet.py index cef6902d..dd92d913 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -243,9 +243,7 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int all_mint_urls = list({k.mint_url for k in wallet.keysets.values()}) proof_summary = { - f"{k.mint_url}/{k.unit.name}": sum( - p.amount for p in wallet.proofs if p.id == k.id - ) + f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id) for k in wallet.keysets.values() } # Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet @@ -600,16 +598,11 @@ async def swap_to_primary_mint( # advance the counter so the next request derives fresh secrets. logger.warning( "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", - extra={ - "mint_quote_id": mint_quote.quote, - "minted_amount": minted_amount, - }, + extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, ) try: for keyset_id in primary_wallet.keysets: - await primary_wallet.restore_tokens_for_keyset( - keyset_id, to=1, batch=25 - ) + await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) await primary_wallet.load_proofs(reload=True) post_recovery_balance = primary_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance From 16fc548b48d982211f4f4a8afde63a6a36f6da9c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 20:12:55 +0200 Subject: [PATCH 33/46] fix: make model path identity selectable --- docs/api/endpoints.md | 56 +-- .../4e0c3d195a49_add_model_paths_table.py | 4 + routstr/core/admin.py | 28 ++ routstr/core/db.py | 14 +- routstr/payment/models.py | 10 +- routstr/upstream/model_paths.py | 107 +++--- tests/unit/test_model_paths.py | 347 +++++++++++++----- 7 files changed, 401 insertions(+), 165 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index 44e6207e..2d202659 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -329,8 +329,8 @@ GET /v1/models ### List Model Paths -Get the upstream provider paths each advertised model can be reached through. -This is discovery data only; routing still chooses the provider per request. +Get the selectable upstream routes for each advertised model. This endpoint is +discovery-only; request-side selection will be added separately. ```http GET /v1/models/paths @@ -342,45 +342,45 @@ GET /v1/models/paths { "data": [ { - "id": "claude-sonnet-4", + "id": "anthropic/claude-sonnet-4", "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"} + { + "path": "provider=12", + "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, + "endpoint": null + }, + { + "path": "provider=42&endpoint=google-vertex%2Fus", + "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, + "endpoint": {"tag": "google-vertex/us", "name": "Google"} + } ] } - ] + ], + "updated_at": 1753500000 } ``` +`path` is an opaque, percent-encoded selector. Clients must store and return it +unchanged rather than parsing or reconstructing it. The configured provider's +stable node-local ID 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 +not silently fall back to another backend. + ### List Paths for One Model -Use a query parameter so model IDs containing `/` are handled safely. Lookup is -by the public, unqualified model ID: `glm-5v-turbo` resolves -`z-ai/glm-5v-turbo`, and `deepseek-v4-pro` and `deepseek/deepseek-v4-pro` -return the same merged path set. +Use the exact model ID advertised by `/v1/models`. The query parameter safely +supports IDs containing `/`. ```http GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 ``` -**Response:** - -```json -{ - "data": [ - {"path": "anthropic"}, - {"path": "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). +The response uses the same path objects and `updated_at` field as the collection +endpoint. An unknown model returns `404 Model not found`. A known model whose +paths have not been discovered yet returns `200` with an empty `data` array. ## Wallet Management diff --git a/migrations/versions/4e0c3d195a49_add_model_paths_table.py b/migrations/versions/4e0c3d195a49_add_model_paths_table.py index 5641e488..688e2bc5 100644 --- a/migrations/versions/4e0c3d195a49_add_model_paths_table.py +++ b/migrations/versions/4e0c3d195a49_add_model_paths_table.py @@ -22,6 +22,10 @@ def upgrade() -> None: sa.Column("id", sa.Integer(), nullable=False), sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("provider_slug", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("provider_type", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("endpoint_tag", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("endpoint_name", sqlmodel.sql.sqltypes.AutoString(), nullable=True), sa.Column("upstream_provider_id", sa.Integer(), nullable=False), sa.Column("updated_at", sa.Integer(), nullable=False, server_default="0"), sa.ForeignKeyConstraint( diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 66a1d288..f581b704 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -51,6 +51,23 @@ ADMIN_SESSION_DURATION = 3600 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 + + 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__, + }, + ) + + async def require_admin_api(request: Request) -> None: auth_header = request.headers.get("Authorization") if not auth_header or not auth_header.startswith("Bearer "): @@ -579,6 +596,7 @@ async def upsert_provider_model( await session.refresh(row) await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return _row_to_model( row, apply_provider_fee=True, provider_fee=provider.provider_fee ).dict() # type: ignore @@ -633,6 +651,7 @@ async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, ob await session.delete(row) await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return {"ok": True, "deleted_id": model_id} @@ -652,6 +671,7 @@ async def delete_all_provider_models(provider_id: str) -> dict[str, object]: await session.delete(row) # type: ignore await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return {"ok": True, "deleted": len(rows)} @@ -705,6 +725,9 @@ 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 @@ -735,6 +758,7 @@ 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) @@ -743,6 +767,7 @@ async def batch_override_provider_models( await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return { "ok": True, "count": overridden_count, @@ -943,6 +968,7 @@ async def create_upstream_provider( await reinitialize_upstreams() await refresh_model_maps() + await _refresh_provider_model_paths(_provider_pk(provider)) return _serialize_provider(provider) @@ -968,6 +994,7 @@ async def update_upstream_provider( await reinitialize_upstreams() await refresh_model_maps() + await _refresh_provider_model_paths(_provider_pk(provider)) return _serialize_provider(provider) @@ -1003,6 +1030,7 @@ async def update_upstream_provider_by_slug( await reinitialize_upstreams() await refresh_model_maps() + await _refresh_provider_model_paths(_provider_pk(provider)) return _serialize_provider(provider) diff --git a/routstr/core/db.py b/routstr/core/db.py index 8b3586ec..151d8627 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -338,8 +338,18 @@ class ModelPathRow(SQLModel, table=True): # type: ignore description="Client-visible /v1/models id (forwarded_model_id or id)" ) path: str = Field( - description="Provider path stamped on chat completion responses, e.g. " - "'anthropic' or 'openrouter:Anthropic'" + description="Opaque selector containing provider ID and optional endpoint tag" + ) + provider_slug: str = Field( + description="Public slug of the configured upstream provider" + ) + provider_type: str = Field(description="Configured upstream provider type") + endpoint_tag: str | None = Field( + default=None, + description="Exact OpenRouter endpoint tag used for request-side selection", + ) + endpoint_name: str | None = Field( + default=None, description="Human-readable endpoint display name" ) upstream_provider_id: int = Field( index=True, diff --git a/routstr/payment/models.py b/routstr/payment/models.py index c433ddfa..5c634ced 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -613,9 +613,17 @@ async def model_paths_for_model(model_id: str) -> dict: model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL encoding and there is no dynamic-route ambiguity. """ + from ..proxy import get_unique_models from ..upstream.model_paths import get_paths_for_model - return await get_paths_for_model(model_id) + result = await get_paths_for_model(model_id) + if not result["data"]: + advertised_ids = { + model.forwarded_model_id or model.id for model in get_unique_models() + } + if model_id not in advertised_ids: + raise HTTPException(status_code=404, detail="Model not found") + return result @models_router.get("/v1/models") diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index f3bb9971..45ce6887 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -4,12 +4,13 @@ Exposes every selectable upstream route a Routstr model is reachable through. This PR remains discovery-only: request-side routing will consume the opaque selectors in a follow-up. -A path is a standard percent-encoded query string containing the normalized -upstream URL and, for an exact OpenRouter endpoint, its machine-readable tag. -Display names never participate in identity:: +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 +machine-readable tag. Upstream URLs and display names never participate in or +leak through public identity:: - url=https%3A%2F%2Fapi.anthropic.com%2Fv1 - url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider=google-vertex%2Fus + provider=42 + provider=42&endpoint=google-vertex%2Fus """ from __future__ import annotations @@ -57,15 +58,24 @@ class EndpointIdentity: provider_name: str | None +@dataclass(frozen=True) +class ConfiguredProviderIdentity: + """Public-safe identity of one configured upstream provider.""" + + id: int + slug: str + provider_type: str + + @dataclass(frozen=True) class DiscoveredPath: """One model route ready for persistence and API serialization.""" model_id: str path: str - upstream_url: str - provider_tag: str | None = None - provider_name: str | None = None + provider: ConfiguredProviderIdentity + endpoint_tag: str | None = None + endpoint_name: str | None = None @dataclass(frozen=True) @@ -76,16 +86,11 @@ class ProviderPathSnapshot: preserve_model_ids: frozenset[str] = frozenset() -def normalize_upstream_url(base_url: str) -> str: - """Normalize route identity without changing URL semantics.""" - return base_url.rstrip("/") - - -def encode_model_path(base_url: str, provider_tag: str | None = None) -> str: - """Encode a stable opaque selector for future request-side routing.""" - components = [("url", normalize_upstream_url(base_url))] - if provider_tag: - components.append(("provider", provider_tag)) +def encode_model_path(provider_id: int, endpoint_tag: str | None = None) -> str: + """Encode a stable opaque selector without exposing upstream URLs.""" + components: list[tuple[str, str | int]] = [("provider", provider_id)] + if endpoint_tag: + components.append(("endpoint", endpoint_tag)) return urlencode(components) @@ -218,9 +223,11 @@ async def _fetch_openrouter_endpoint_subproviders( return None try: - endpoints = resp.json().get("data", {}).get("endpoints", []) + payload = resp.json() + data = payload.get("data") if isinstance(payload, dict) else None + endpoints = data.get("endpoints") if isinstance(data, dict) else None if not isinstance(endpoints, list): - endpoints = [] + raise ValueError("endpoints must be a list") identities: dict[str, EndpointIdentity] = {} for endpoint in endpoints: if not isinstance(endpoint, dict): @@ -238,6 +245,8 @@ async def _fetch_openrouter_endpoint_subproviders( else None, ), ) + if endpoints and not identities: + raise ValueError("endpoints contain no usable tags") result = list(identities.values()) except Exception as e: # noqa: BLE001 logger.warning( @@ -251,7 +260,9 @@ async def _fetch_openrouter_endpoint_subproviders( async def _load_model_visibility() -> tuple[ - dict[ModelKey, ModelRow], set[ModelKey], set[int] + dict[ModelKey, ModelRow], + set[ModelKey], + dict[int, ConfiguredProviderIdentity], ]: """Load the same DB model visibility inputs used by routing. @@ -271,12 +282,16 @@ async def _load_model_visibility() -> tuple[ overrides_by_key: dict[ModelKey, ModelRow] = {} disabled_model_keys: set[ModelKey] = set() - enabled_provider_ids: set[int] = set() + provider_identities: dict[int, ConfiguredProviderIdentity] = {} for provider in provider_rows: if not provider.enabled or provider.id is None: continue - enabled_provider_ids.add(provider.id) + provider_identities[provider.id] = ConfiguredProviderIdentity( + id=provider.id, + slug=provider.slug or f"provider-{provider.id}", + provider_type=provider.provider_type, + ) for model in provider.models: key = (model.id.lower(), provider.id) if model.enabled: @@ -284,7 +299,7 @@ async def _load_model_visibility() -> tuple[ else: disabled_model_keys.add(key) - return overrides_by_key, disabled_model_keys, enabled_provider_ids + return overrides_by_key, disabled_model_keys, provider_identities def _apply_model_visibility( @@ -338,6 +353,7 @@ def _apply_model_visibility( async def _collect_provider_paths( upstream: BaseUpstreamProvider, + provider_identity: ConfiguredProviderIdentity, overrides_by_key: dict[ModelKey, ModelRow] | None = None, disabled_model_keys: set[ModelKey] | None = None, cycle: _RefreshCycleState | None = None, @@ -350,13 +366,12 @@ async def _collect_provider_paths( """ cycle = cycle or _RefreshCycleState() models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) - upstream_url = normalize_upstream_url(upstream.base_url) def _base_path(model: object) -> DiscoveredPath: return DiscoveredPath( model_id=exposed_model_id(model), - path=encode_model_path(upstream_url), - upstream_url=upstream_url, + path=encode_model_path(provider_identity.id), + provider=provider_identity, ) if not is_openrouter_base_url(upstream.base_url): @@ -389,10 +404,10 @@ async def _collect_provider_paths( paths.extend( DiscoveredPath( model_id=model_id, - path=encode_model_path(upstream_url, endpoint.tag), - upstream_url=upstream_url, - provider_tag=endpoint.tag, - provider_name=endpoint.provider_name, + path=encode_model_path(provider_identity.id, endpoint.tag), + provider=provider_identity, + endpoint_tag=endpoint.tag, + endpoint_name=endpoint.provider_name, ) for endpoint in endpoints ) @@ -448,9 +463,10 @@ async def _persist_provider_paths( { "model_id": discovered.model_id, "path": discovered.path, - "upstream_url": discovered.upstream_url, - "provider_tag": discovered.provider_tag, - "provider_name": discovered.provider_name, + "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, } @@ -505,17 +521,18 @@ async def refresh_model_paths( ( overrides_by_key, disabled_model_keys, - enabled_provider_ids, + provider_identities, ) = await _load_model_visibility() await prune_model_paths_for_inactive_providers() cycle = _RefreshCycleState() for upstream in upstreams: - if upstream.db_id is None or upstream.db_id not in enabled_provider_ids: + if upstream.db_id is None or upstream.db_id not in provider_identities: continue try: snapshot = await _collect_provider_paths( upstream, + provider_identity=provider_identities[upstream.db_id], overrides_by_key=overrides_by_key, disabled_model_keys=disabled_model_keys, cycle=cycle, @@ -611,13 +628,17 @@ async def refresh_model_paths_periodically( def _serialize_path(row: ModelPathRow) -> dict[str, Any]: - provider = None - if row.provider_tag or row.provider_name: - provider = {"name": row.provider_name, "slug": row.provider_tag} + endpoint = None + if row.endpoint_tag or row.endpoint_name: + endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name} return { "path": row.path, - "upstream_url": row.upstream_url, - "provider": provider, + "provider": { + "id": row.upstream_provider_id, + "slug": row.provider_slug, + "type": row.provider_type, + }, + "endpoint": endpoint, } @@ -646,9 +667,7 @@ async def get_all_model_paths() -> dict: data = [ { "id": grouped_model_id, - "paths": sorted( - grouped[grouped_model_id], key=lambda item: str(item["path"]) - ), + "paths": grouped[grouped_model_id], } for grouped_model_id in sorted(grouped) ] diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 7d3393fe..0538ea55 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -142,11 +142,18 @@ def _mock_transport( return counter -def _endpoints_response(*provider_names: str) -> httpx.Response: - return httpx.Response( - 200, - json={"data": {"endpoints": [{"provider_name": n} for n in provider_names]}}, - ) +def _endpoints_response( + *providers: str | tuple[str, str], +) -> httpx.Response: + endpoints = [] + for provider in providers: + if isinstance(provider, tuple): + provider_name, tag = provider + else: + provider_name = provider + tag = provider.lower().replace(" ", "-") + endpoints.append({"provider_name": provider_name, "tag": tag}) + return httpx.Response(200, json={"data": {"endpoints": endpoints}}) _SEEDED_PROVIDER_IDS = (1, 2, 4, 5, 7) @@ -206,6 +213,29 @@ def _ids_of(payload: dict) -> set[str]: return {entry["id"] for entry in payload["data"]} +def _path_entry( + provider_id: int, + *, + provider_slug: str | None = None, + provider_type: str | None = None, + endpoint_tag: str | None = None, + endpoint_name: str | None = None, +) -> dict[str, Any]: + endpoint = None + if endpoint_tag or endpoint_name: + endpoint = {"tag": endpoint_tag, "name": endpoint_name} + return { + "path": mp.encode_model_path(provider_id, endpoint_tag), + "provider": { + "id": provider_id, + "slug": provider_slug or f"p{provider_id}", + "type": provider_type + or ("anthropic" if provider_id == 1 else "openrouter"), + }, + "endpoint": endpoint, + } + + # --------------------------------------------------------------------------- # # Predicates / pure helpers # --------------------------------------------------------------------------- # @@ -223,6 +253,13 @@ 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_exposed_model_id_prefers_forwarded() -> None: assert ( mp.exposed_model_id(_model("claude-x", forwarded_model_id="fwd-claude")) @@ -300,9 +337,7 @@ async def test_direct_provider_single_path_uses_provider_type( ) await mp.refresh_model_paths([provider]) payload = await mp.get_all_model_paths() - assert payload["data"] == [ - {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} - ] + assert payload["data"] == [{"id": "claude-opus-4.6", "paths": [_path_entry(1)]}] assert payload["updated_at"] is not None @@ -320,6 +355,22 @@ async def test_direct_path_stores_exposed_model_id( assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} +@pytest.mark.asyncio +async def test_forwarded_model_id_with_slash_remains_exact_and_routable( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("local-alias", forwarded_model_id="anthropic/claude-opus-4.6")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _ids_of(await mp.get_all_model_paths()) == {"anthropic/claude-opus-4.6"} + assert (await mp.get_paths_for_model("anthropic/claude-opus-4.6"))["data"] + + @pytest.mark.asyncio async def test_disabled_cached_models_excluded( patched_session: AsyncEngine, @@ -361,7 +412,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") == {"anthropic"} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} @pytest.mark.asyncio @@ -395,8 +446,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") == {"anthropic"} - assert _paths_of(payload, "private-alias") == {"generic"} + assert _paths_of(payload, "shared-model") == {mp.encode_model_path(1)} + assert _paths_of(payload, "private-alias") == {mp.encode_model_path(2)} @pytest.mark.asyncio @@ -464,7 +515,7 @@ async def test_refresh_model_paths_uses_db_forwarded_alias( await mp.refresh_model_paths([provider]) assert (await mp.get_all_model_paths())["data"] == [ - {"id": "public-alias", "paths": [{"path": "anthropic"}]} + {"id": "public-alias", "paths": [_path_entry(1)]} ] @@ -486,7 +537,7 @@ async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cac await mp.refresh_model_paths([provider]) assert (await mp.get_all_model_paths())["data"] == [ - {"id": "public-deployment", "paths": [{"path": "generic"}]} + {"id": "public-deployment", "paths": [_path_entry(1)]} ] @@ -638,24 +689,34 @@ async def test_openrouter_provider_adds_endpoint_paths( ) _mock_transport( monkeypatch, - lambda request: _endpoints_response("Anthropic", "Amazon Bedrock"), + lambda request: _endpoints_response( + ("Google", "google-vertex/eu"), + ("Google", "google-vertex/us"), + ), ) await mp.refresh_model_paths([provider]) - paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert "openrouter:Anthropic" in paths - assert "openrouter:Amazon Bedrock" in paths - assert "openrouter" not in paths + 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"), + } + assert { + item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] + } == {"google-vertex/eu", "google-vertex/us"} + assert { + item["endpoint"]["name"] for item in payload["data"] if item["endpoint"] + } == {"Google"} + assert {item["provider"]["id"] for item in payload["data"]} == {2} @pytest.mark.asyncio -async def test_openrouter_self_echoing_subprovider_maps_to_unknown( +async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch ) -> None: - """Responses stamped for a sub-provider echoing "OpenRouter" say - ``unknown``; discovery must advertise the same string, never - ``openrouter:OpenRouter``.""" + """Machine-readable endpoint tags, not display names, define identity.""" provider = _FakeOpenRouterProvider( models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], db_id=2, @@ -665,17 +726,14 @@ async def test_openrouter_self_echoing_subprovider_maps_to_unknown( await mp.refresh_model_paths([provider]) paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") - assert "openrouter:OpenRouter" not in paths - assert "unknown" in paths + assert paths == {mp.encode_model_path(2), mp.encode_model_path(2, "openrouter")} @pytest.mark.asyncio async def test_generic_provider_with_openrouter_base_url_discovers( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch ) -> None: - """A generic provider pointed at OpenRouter exposes both the bare - ``generic`` path (stamped when the upstream omits its provider field) and - the ``generic:`` endpoint paths.""" + """Configured provider identity is independent from its endpoint URL.""" provider = _FakeProvider( provider_type="generic", base_url="https://openrouter.ai/api/v1", @@ -687,7 +745,35 @@ 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 == {"generic", "generic:Anthropic"} + assert paths == {mp.encode_model_path(1), mp.encode_model_path(1, "anthropic")} + + +@pytest.mark.asyncio +async def test_openrouter_partial_failure_keeps_failed_models_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + provider = _FakeOpenRouterProvider( + models=[ + _model("good", canonical_slug="author/good"), + _model("degraded", canonical_slug="author/degraded"), + ], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "degraded") + assert before + + def _partial_failure(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/author/degraded/endpoints"): + return httpx.Response(503) + return _endpoints_response("Google") + + _mock_transport(monkeypatch, _partial_failure) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "degraded") == before + assert _paths_of(await mp.get_all_model_paths(), "good") != before @pytest.mark.asyncio @@ -703,7 +789,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 "openrouter:Anthropic" in before + assert mp.encode_model_path(2, "anthropic") in before def _network_down(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("network down", request=request) @@ -726,10 +812,8 @@ 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]) - assert _paths_of(await mp.get_all_model_paths(), "m0") == { - "unknown", - "openrouter:Anthropic", - } + expected = {mp.encode_model_path(2), mp.encode_model_path(2, "anthropic")} + assert _paths_of(await mp.get_all_model_paths(), "m0") == expected counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) await mp.refresh_model_paths([provider]) @@ -739,27 +823,28 @@ async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( assert counter["requests"] <= mp._OPENROUTER_CONCURRENCY, ( "429 must abort the remaining fan-out" ) - assert _paths_of(await mp.get_all_model_paths(), "m0") == { - "unknown", - "openrouter:Anthropic", - } + assert _paths_of(await mp.get_all_model_paths(), "m0") == expected @pytest.mark.asyncio -async def test_openrouter_bad_payload_shapes_do_not_raise( +async def test_openrouter_bad_payload_shapes_preserve_previous_rows( patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch ) -> None: - """``endpoints: null`` and non-list endpoint payloads are swallowed as - documented, not raised into the generic task-errored bucket.""" + """Malformed successful responses are degraded snapshots, not empty sets.""" + provider = _FakeOpenRouterProvider( + models=[_model("m", canonical_slug="a/m")], db_id=2 + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "m") + for payload in ( {"data": {"endpoints": None}}, {"data": {"endpoints": "none"}}, + {"data": {"endpoints": [{"provider_name": "Anthropic"}]}}, {"data": None}, {}, ): - provider = _FakeOpenRouterProvider( - models=[_model("m", canonical_slug="a/m")], db_id=2 - ) def _handler( request: httpx.Request, p: dict[str, Any] | None = payload @@ -767,8 +852,8 @@ async def test_openrouter_bad_payload_shapes_do_not_raise( return httpx.Response(200, json=p) _mock_transport(monkeypatch, _handler) - # Must not raise. await mp.refresh_model_paths([provider]) + assert _paths_of(await mp.get_all_model_paths(), "m") == before @pytest.mark.asyncio @@ -795,8 +880,12 @@ 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 "openrouter:Anthropic" in paths - assert "generic:Anthropic" in paths + assert paths == { + mp.encode_model_path(2), + mp.encode_model_path(2, "anthropic"), + mp.encode_model_path(4), + mp.encode_model_path(4, "anthropic"), + } @pytest.mark.asyncio @@ -857,13 +946,16 @@ async def test_same_model_two_providers_two_paths( assert len(payload["data"]) == 1 entry = payload["data"][0] assert entry["id"] == "claude-opus-4.6" - assert {p["path"] for p in entry["paths"]} == {"anthropic", "generic"} + assert {p["path"] for p in entry["paths"]} == { + mp.encode_model_path(1), + mp.encode_model_path(2), + } assert "canonical_id" not in entry assert all("canonical_id" not in p for p in entry["paths"]) @pytest.mark.asyncio -async def test_get_all_model_paths_deduplicates_visible_paths( +async def test_get_all_model_paths_keeps_distinct_configured_providers( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -881,7 +973,10 @@ async def test_get_all_model_paths_deduplicates_visible_paths( await mp.refresh_model_paths([p1, p2]) assert (await mp.get_all_model_paths())["data"] == [ - {"id": "claude-opus-4.6", "paths": [{"path": "anthropic"}]} + { + "id": "claude-opus-4.6", + "paths": [_path_entry(1), _path_entry(2)], + } ] @@ -906,14 +1001,13 @@ async def test_get_all_model_paths_is_deterministic( @pytest.mark.asyncio -async def test_get_paths_for_model_returns_only_paths( +async def test_get_paths_for_model_returns_route_identity( patched_session: AsyncEngine, ) -> None: await _seed_two_provider_shared_model(patched_session) payload = await mp.get_paths_for_model("claude-opus-4.6") - assert {p["path"] for p in payload["data"]} == {"anthropic", "generic"} - assert all(set(p.keys()) == {"path"} for p in payload["data"]) + assert payload["data"] == [_path_entry(1), _path_entry(2)] assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] @@ -929,13 +1023,11 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( ) await mp.refresh_model_paths([provider]) - assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ - {"path": "generic"} - ] + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [_path_entry(4)] @pytest.mark.asyncio -async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( +async def test_get_paths_for_model_requires_exact_advertised_id( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -955,8 +1047,8 @@ async def test_get_paths_for_model_merges_prefixed_and_unprefixed_aliases( short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] - assert {p["path"] for p in short_paths} == {"generic", "anthropic"} - assert prefixed_paths == short_paths + assert short_paths == [_path_entry(4), _path_entry(7)] + assert prefixed_paths == [] @pytest.mark.asyncio @@ -975,18 +1067,40 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ - {"path": "generic"} + _path_entry(1) ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" - ] == [{"path": "generic"}] + ] == [] # --------------------------------------------------------------------------- # -# Periodic refresh loop +# Immediate and periodic refresh # --------------------------------------------------------------------------- # +@pytest.mark.asyncio +async def test_refresh_model_paths_for_provider_selects_mutated_provider( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + target = SimpleNamespace(db_id=2) + other = SimpleNamespace(db_id=1) + seen: list[list[Any]] = [] + + monkeypatch.setattr(proxy, "get_upstreams", lambda: [other, target]) + + async def _fake_refresh(upstreams: list[Any]) -> None: + seen.append(upstreams) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + await mp.refresh_model_paths_for_provider(2) + + assert seen == [[target]] + + @pytest.mark.asyncio async def test_refresh_loop_rereads_interval_and_picks_up_providers( monkeypatch: pytest.MonkeyPatch, @@ -1093,37 +1207,88 @@ def _make_model_paths_app() -> FastAPI: def test_model_paths_endpoint_returns_all_paths( monkeypatch: pytest.MonkeyPatch, ) -> None: + expected = { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + _path_entry(1), + _path_entry( + 2, + endpoint_tag="google-vertex/us", + endpoint_name="Google", + ), + ], + } + ], + "updated_at": 1753500000, + } + async def _fake_get_all_model_paths() -> dict[str, Any]: - return { - "data": [ - { - "id": "claude-opus-4.6", - "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"}, - ], - } - ], - "updated_at": 1753500000, - } + return expected monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) response = TestClient(_make_model_paths_app()).get("/v1/models/paths") assert response.status_code == 200 - assert response.json() == { - "data": [ - { - "id": "claude-opus-4.6", - "paths": [ - {"path": "anthropic"}, - {"path": "openrouter:Anthropic"}, - ], - } - ], - "updated_at": 1753500000, - } + assert response.json() == expected + + +def test_model_paths_for_model_returns_404_for_unknown_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: []) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "does-not-exist"} + ) + + assert response.status_code == 404 + assert response.json() == {"detail": "Model not found"} + + +def test_model_paths_for_known_model_can_return_empty_collection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: [_model("known")]) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "known"} + ) + + assert response.status_code == 200 + assert response.json() == {"data": [], "updated_at": None} + + +def test_model_paths_for_routing_only_alias_returns_404( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: [_model("advertised")]) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "routing-alias"} + ) + + assert response.status_code == 404 def test_model_paths_for_model_endpoint_accepts_slash_model_id( @@ -1131,9 +1296,14 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) -> None: calls: list[str] = [] + expected = { + "data": [_path_entry(2, endpoint_tag="anthropic", endpoint_name="Anthropic")], + "updated_at": None, + } + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: calls.append(model_id) - return {"data": [{"path": "generic:Anthropic"}], "updated_at": None} + return expected monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) @@ -1143,8 +1313,5 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( ) assert response.status_code == 200 - assert response.json() == { - "data": [{"path": "generic:Anthropic"}], - "updated_at": None, - } + assert response.json() == expected assert calls == ["anthropic/claude-opus-4.6"] From c829685f806d2397851546b9a98877806550a6d0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 23:23:32 +0200 Subject: [PATCH 34/46] fix Cashu fallback and Lightning settlement --- routstr/lightning.py | 295 ++++++++++++++---- routstr/payment/helpers.py | 7 +- routstr/proxy.py | 5 +- routstr/wallet.py | 50 ++- .../integration/test_insufficient_balance.py | 10 +- .../test_lightning_invoice_constraints.py | 11 +- .../integration/test_lightning_settlement.py | 278 +++++++++++++++++ tests/unit/test_lightning_settlement.py | 206 ++++++++++++ tests/unit/test_payment_helpers.py | 6 +- tests/unit/test_stale_reservations.py | 2 +- tests/unit/test_upstream_rate_limit.py | 2 +- tests/unit/test_wallet.py | 69 ++-- 12 files changed, 825 insertions(+), 116 deletions(-) create mode 100644 tests/integration/test_lightning_settlement.py create mode 100644 tests/unit/test_lightning_settlement.py diff --git a/routstr/lightning.py b/routstr/lightning.py index c4bc2b20..43d1a519 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -1,11 +1,14 @@ import asyncio import hashlib +import re import secrets import time +from dataclasses import dataclass +from typing import Any from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel, Field -from sqlmodel import col, select +from sqlmodel import col, select, update from sqlmodel.ext.asyncio.session import AsyncSession from .core.db import ApiKey, LightningInvoice, create_session, get_session @@ -24,6 +27,37 @@ logger = get_logger(__name__) lightning_router = APIRouter(prefix="/lightning") +# Avoid duplicate work within one process. Cross-process credit fencing is done +# by the conditional pending -> paid update in _finalize_invoice_settlement(). +_invoice_settlement_locks: dict[str, asyncio.Lock] = {} + + +@dataclass(frozen=True) +class _InvoiceSettlement: + id: str + payment_hash: str + amount_sats: int + purpose: str + api_key_hash: str | None + mint_url: str | None + balance_limit: int | None + balance_limit_reset: str | None + validity_date: int | None + + @classmethod + def from_invoice(cls, invoice: LightningInvoice) -> "_InvoiceSettlement": + return cls( + id=invoice.id, + payment_hash=invoice.payment_hash, + amount_sats=invoice.amount_sats, + purpose=invoice.purpose, + api_key_hash=invoice.api_key_hash, + mint_url=invoice.mint_url, + balance_limit=invoice.balance_limit, + balance_limit_reset=invoice.balance_limit_reset, + validity_date=invoice.validity_date, + ) + class InvoiceCreateRequest(BaseModel): amount_sats: int = Field(gt=0, le=1_000_000, description="Amount in satoshis") @@ -306,42 +340,203 @@ async def recover_invoice( async def check_invoice_payment( invoice: LightningInvoice, session: AsyncSession ) -> None: - try: - mint_url = invoice.mint_url or settings.primary_mint - wallet = await get_wallet(mint_url, "sat") - - mint_status = await _mint_operation( - lambda: wallet.get_mint_quote(invoice.payment_hash), - op_name="get_mint_quote", - mint_url=mint_url, - ) - - if mint_status.paid: - invoice.status = "paid" - invoice.paid_at = int(time.time()) - - if invoice.purpose == "create": - api_key = await create_api_key_from_invoice(invoice, session) - invoice.api_key_hash = api_key.hashed_key - elif invoice.purpose == "topup" and invoice.api_key_hash: - await topup_api_key_from_invoice(invoice, session) - + lock = _invoice_settlement_locks.setdefault(invoice.id, asyncio.Lock()) + async with lock: + try: + # Refresh and snapshot the row, then close the read transaction before + # any wallet or mint network I/O. The final DB mutations use a new, + # short transaction and a conditional status update as their fence. + await session.refresh(invoice) + if invoice.status != "pending": + await session.commit() + return + settlement = _InvoiceSettlement.from_invoice(invoice) await session.commit() + mint_url = settlement.mint_url or settings.primary_mint + wallet = await get_wallet(mint_url, "sat") + mint_status = await _mint_operation( + lambda: wallet.get_mint_quote(settlement.payment_hash), + op_name="get_mint_quote", + mint_url=mint_url, + ) + if not mint_status.paid: + return + + await _mint_invoice_quote(wallet, settlement) + paid_at = int(time.time()) + settled, api_key_hash = await _finalize_invoice_settlement( + settlement, session, paid_at + ) + if not settled: + await _reload_invoice_view(invoice, session) + return + + invoice.status = "paid" + invoice.paid_at = paid_at + invoice.api_key_hash = api_key_hash logger.info( "Lightning invoice paid", extra={ - "invoice_id": invoice.id, - "amount_sats": invoice.amount_sats, - "purpose": invoice.purpose, - "api_key_hash": invoice.api_key_hash[:8] + "..." - if invoice.api_key_hash + "invoice_id": settlement.id, + "amount_sats": settlement.amount_sats, + "purpose": settlement.purpose, + "api_key_hash": api_key_hash[:8] + "..." + if api_key_hash else None, }, ) - except Exception as e: - logger.error(f"Failed to check invoice payment: {e}") + except Exception as e: + await session.rollback() + try: + await _reload_invoice_view(invoice, session) + except Exception: + pass + logger.error(f"Failed to check invoice payment: {e}") + + +def _is_outputs_already_signed(error: BaseException) -> bool: + message = str(error) + return bool( + re.search( + r"\boutputs?\s+(?:have\s+)?already\s+(?:been\s+)?signed(?:\s+before)?\b", + message, + re.IGNORECASE, + ) + and re.search(r"\bcode\s*:\s*11003\b", message, re.IGNORECASE) + ) + + +def _invoice_quote_proof_amount(wallet: Any, quote_id: str) -> int: + """Return spendable wallet value minted by one Lightning quote.""" + return sum( + proof.amount + for proof in wallet.proofs + if proof.mint_id == quote_id and not proof.reserved + ) + + +async def _mint_invoice_quote( + wallet: Any, invoice: LightningInvoice | _InvoiceSettlement +) -> None: + """Mint a paid quote, proving quote-linked outputs before DB credit.""" + mint_url = invoice.mint_url or settings.primary_mint + await wallet.load_proofs(reload=True) + if _invoice_quote_proof_amount(wallet, invoice.payment_hash) >= invoice.amount_sats: + return + + try: + await _mint_operation( + lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), + op_name=f"invoice_mint_{invoice.purpose}", + mint_url=mint_url, + retry_timeouts=False, + ) + except Exception as error: + if not _is_outputs_already_signed(error): + raise + + for keyset_id in wallet.keysets: + await wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25) + await wallet.load_proofs(reload=True) + recovered = _invoice_quote_proof_amount(wallet, invoice.payment_hash) + if recovered < invoice.amount_sats: + raise RuntimeError( + "Invoice outputs were already signed but quote-linked recovery returned " + f"{recovered} sats; expected at least {invoice.amount_sats}" + ) from error + else: + await wallet.load_proofs(reload=True) + minted = _invoice_quote_proof_amount(wallet, invoice.payment_hash) + if minted < invoice.amount_sats: + raise RuntimeError( + "Invoice mint succeeded but quote-linked proofs total " + f"{minted} sats; expected at least {invoice.amount_sats}" + ) + + +def _invoice_api_key_hash(invoice: LightningInvoice | _InvoiceSettlement) -> str: + dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" + return hashlib.sha256(dummy_token.encode()).hexdigest() + + +async def _create_api_key_record( + invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession +) -> ApiKey: + mint_url = invoice.mint_url or settings.primary_mint + api_key = ApiKey( + hashed_key=_invoice_api_key_hash(invoice), + balance=invoice.amount_sats * 1000, + refund_currency="sat", + refund_mint_url=mint_url, + balance_limit=invoice.balance_limit, + balance_limit_reset=invoice.balance_limit_reset, + validity_date=invoice.validity_date, + ) + session.add(api_key) + await session.flush() + return api_key + + +async def _topup_api_key_record( + invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession +) -> None: + if not invoice.api_key_hash: + raise ValueError("No API key associated with topup invoice") + result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == invoice.api_key_hash) + .values(balance=col(ApiKey.balance) + invoice.amount_sats * 1000) + .execution_options(synchronize_session=False) + ) + if result.rowcount != 1: + raise ValueError("Associated API key not found") + + +async def _finalize_invoice_settlement( + invoice: _InvoiceSettlement, session: AsyncSession, paid_at: int +) -> tuple[bool, str | None]: + """Atomically fence and apply one invoice credit across all processes.""" + api_key_hash = ( + _invoice_api_key_hash(invoice) + if invoice.purpose == "create" + else invoice.api_key_hash + ) + claim = await session.exec( # type: ignore[call-overload] + update(LightningInvoice) + .where(col(LightningInvoice.id) == invoice.id) + .where(col(LightningInvoice.status) == "pending") + .values(status="paid", paid_at=paid_at, api_key_hash=api_key_hash) + .execution_options(synchronize_session=False) + ) + if claim.rowcount != 1: + await session.rollback() + return False, None + + try: + if invoice.purpose == "create": + await _create_api_key_record(invoice, session) + elif invoice.purpose == "topup": + await _topup_api_key_record(invoice, session) + else: + raise ValueError(f"Unsupported invoice purpose: {invoice.purpose}") + await session.commit() + except Exception: + await session.rollback() + raise + return True, api_key_hash + + +async def _reload_invoice_view( + invoice: LightningInvoice, session: AsyncSession +) -> None: + stored = await session.get(LightningInvoice, invoice.id) + if stored is not None: + invoice.status = stored.status + invoice.paid_at = stored.paid_at + invoice.api_key_hash = stored.api_key_hash + await session.commit() async def create_api_key_from_invoice( @@ -349,30 +544,8 @@ async def create_api_key_from_invoice( ) -> ApiKey: mint_url = invoice.mint_url or settings.primary_mint wallet = await get_wallet(mint_url, "sat") - await _mint_operation( - lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), - op_name="invoice_mint_create", - mint_url=mint_url, - retry_timeouts=False, - ) - - dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" - hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest() - - api_key = ApiKey( - hashed_key=hashed_key, - balance=invoice.amount_sats * 1000, # Convert to msats - refund_currency="sat", - refund_mint_url=mint_url, - balance_limit=invoice.balance_limit, - balance_limit_reset=invoice.balance_limit_reset, - validity_date=invoice.validity_date, - ) - - session.add(api_key) - await session.flush() - - return api_key + await _mint_invoice_quote(wallet, invoice) + return await _create_api_key_record(invoice, session) async def topup_api_key_from_invoice( @@ -380,22 +553,8 @@ async def topup_api_key_from_invoice( ) -> None: mint_url = invoice.mint_url or settings.primary_mint wallet = await get_wallet(mint_url, "sat") - await _mint_operation( - lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), - op_name="invoice_mint_topup", - mint_url=mint_url, - retry_timeouts=False, - ) - - if not invoice.api_key_hash: - raise ValueError("No API key associated with topup invoice") - - api_key = await session.get(ApiKey, invoice.api_key_hash) - if not api_key: - raise ValueError("Associated API key not found") - - api_key.balance += invoice.amount_sats * 1000 # Convert to msats - await session.flush() + await _mint_invoice_quote(wallet, invoice) + await _topup_api_key_record(invoice, session) # Nutshell mints throttle Lightning backend lookups to once per 10s per diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 592bda33..67a7284d 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -18,11 +18,14 @@ from ..wallet import deserialize_token_from_string logger = get_logger(__name__) -_MINT_FEE_ALLOWANCE = 0.10 +# Interim policy: when Routstr must move value to another trusted mint, the +# cross-mint Lightning round trip can consume fees that are not visible to the +# client. Reserve 5% headroom until the fee-payer policy is made explicit. +_MINT_FEE_ALLOWANCE = 0.05 def apply_mint_fee_allowance(cost_msat: int) -> int: - """Reduce the admission reservation to account for mint fallback fees.""" + """Reserve headroom for possible trusted-mint fallback fees.""" adjusted = math.ceil(cost_msat * (1 - _MINT_FEE_ALLOWANCE)) return max(settings.min_request_msat, adjusted) diff --git a/routstr/proxy.py b/routstr/proxy.py index 9e66c86b..7331fa3c 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -25,7 +25,6 @@ from .core.db import ( ) from .core.exceptions import UpstreamError from .core.not_found import build_not_found_response -from .core.settings import settings from .payment.helpers import ( apply_mint_fee_allowance, calculate_discounted_max_cost, @@ -469,7 +468,9 @@ async def proxy( candidate_max = await calculate_discounted_max_cost( candidate_max, request_body_dict, model_obj=model_obj ) - candidate_max = max(candidate_max, settings.min_request_msat) + # Apply the same interim 5% trusted-mint fee headroom used for the + # first candidate; failover must not silently change admission. + candidate_max = apply_mint_fee_allowance(candidate_max) if candidate_max > max_cost_for_model: await revert_pay_for_request( key, session, max_cost_for_model, reservation_snapshot diff --git a/routstr/wallet.py b/routstr/wallet.py index 3c13bbc5..e72da370 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -9,7 +9,7 @@ import httpx from cashu.core.base import MintQuote, Proof, Token from cashu.core.mint_info import MintInfo as _CashuMintInfo from cashu.wallet.helpers import deserialize_token_from_string -from cashu.wallet.wallet import Wallet +from cashu.wallet.wallet import Wallet as _CashuWallet from pydantic_core import PydanticUndefined from sqlmodel import col, select, update @@ -55,6 +55,30 @@ class TokenConsumedError(Exception): """ +class MintRateLimitedError(httpx.HTTPStatusError): + """Typed boundary error preserving a Cashu mint's HTTP 429 response.""" + + +class Wallet(_CashuWallet): + """Cashu wallet adapter that preserves rate-limit status information. + + Cashu's default response adapter converts JSON error bodies into plain + ``Exception`` instances before calling ``raise_for_status``. Intercept 429 + here so Routstr's fallback and cooldown policy can use the real status + without unreliable message matching. + """ + + @staticmethod + def raise_on_error_request(resp: httpx.Response) -> None: + if resp.status_code == 429: + raise MintRateLimitedError( + "Cashu mint rate limited", + request=resp.request, + response=resp, + ) + _CashuWallet.raise_on_error_request(resp) + + # httpx base classes cover their subclasses. HTTPStatusError is excluded on # purpose — that means the mint answered, just with an error status. _MINT_TRANSPORT_COOLDOWN_SECONDS = 30.0 @@ -253,8 +277,10 @@ async def _mint_operation( ) -> Any: """Run a mint operation with bounded concurrency and adaptive cooldown. - The timeout covers concurrency queueing, 429 cooldown, backoff, and network - work together. ``factory`` must return a fresh coroutine for every retry. + The timeout applies to each network attempt. Queueing, cooldown, and retry + backoff are deliberately outside it so the shipped 60-second 429 cooldown + is not cancelled by the 30-second operation timeout. ``factory`` must + return a fresh coroutine for every retry. When ``retry_on_rate_limit`` is False a 429 is not retried in-place — the cooldown is still applied to the per-mint guard (so subsequent operations on @@ -265,10 +291,15 @@ async def _mint_operation( timeout = settings.mint_operation_timeout_seconds max_attempts = settings.mint_retry_max_attempts + 1 + async def timed_factory() -> Any: + if timeout > 0: + return await asyncio.wait_for(factory(), timeout=timeout) + return await factory() + async def invoke() -> Any: if guard is not None: - return await guard.run(factory) - return await factory() + return await guard.run(timed_factory) + return await timed_factory() async def run_with_retries() -> Any: for attempt in range(max_attempts): @@ -343,14 +374,7 @@ async def _mint_operation( raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") - try: - if timeout > 0: - return await asyncio.wait_for(run_with_retries(), timeout=timeout) - return await run_with_retries() - except asyncio.TimeoutError as exc: - raise httpx.TimeoutException( - f"{op_name} exceeded its {timeout}s total timeout" - ) from exc + return await run_with_retries() def _parse_retry_after(headers: Any) -> float | None: diff --git a/tests/integration/test_insufficient_balance.py b/tests/integration/test_insufficient_balance.py index bcdbfd17..3acc5db9 100644 --- a/tests/integration/test_insufficient_balance.py +++ b/tests/integration/test_insufficient_balance.py @@ -208,13 +208,13 @@ async def test_pay_for_request_succeeds_when_balance_equals_cost( @pytest.mark.asyncio -async def test_ten_percent_mint_fee_shortfall_is_admitted_and_reserved( +async def test_five_percent_mint_fallback_headroom_is_admitted_and_reserved( integration_session: AsyncSession, ) -> None: from routstr.auth import pay_for_request, validate_bearer_key from routstr.payment.helpers import apply_mint_fee_allowance - key = _key(balance=90_000) + key = _key(balance=95_000) integration_session.add(key) await integration_session.commit() @@ -225,8 +225,8 @@ async def test_ten_percent_mint_fee_shortfall_is_admitted_and_reserved( await pay_for_request(validated, admission_cost, integration_session) await integration_session.refresh(key) - assert admission_cost == 90_000 - assert key.reserved_balance == 90_000 + assert admission_cost == 95_000 + assert key.reserved_balance == 95_000 # --------------------------------------------------------------------------- @@ -288,7 +288,7 @@ async def test_http_402_response_shape_on_insufficient_balance( error = body["detail"]["error"] assert error["code"] == "insufficient_balance" assert error["type"] == "insufficient_quota" - assert "560.6 sats (560600 msats) required" in error["message"] + assert "591.744 sats (591744 msats) required" in error["message"] assert "20.32 sats (20320 msats) available" in error["message"] # Balance must be completely untouched diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index a26b9083..413878cf 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -13,6 +13,7 @@ import time from unittest.mock import AsyncMock, patch import pytest +from cashu.core.base import Proof from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.db import ApiKey, LightningInvoice @@ -39,7 +40,15 @@ def _make_invoice(**kwargs: object) -> LightningInvoice: def mock_wallet_mint() -> object: with patch("routstr.lightning.get_wallet") as mock_get_wallet: wallet = AsyncMock() - wallet.mint = AsyncMock(return_value=[]) + wallet.proofs = [] + wallet.load_proofs = AsyncMock() + + async def mint(amount: int, quote_id: str) -> list[Proof]: + proofs = [Proof(amount=amount, mint_id=quote_id)] + wallet.proofs.extend(proofs) + return proofs + + wallet.mint = AsyncMock(side_effect=mint) mock_get_wallet.return_value = wallet yield mock_get_wallet diff --git a/tests/integration/test_lightning_settlement.py b/tests/integration/test_lightning_settlement.py new file mode 100644 index 00000000..dff37c48 --- /dev/null +++ b/tests/integration/test_lightning_settlement.py @@ -0,0 +1,278 @@ +import asyncio +import time +import uuid +from unittest.mock import AsyncMock, Mock, patch + +import pytest +from cashu.core.base import Proof +from sqlmodel import col, update +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey, LightningInvoice +from routstr.lightning import ( + _finalize_invoice_settlement, + _InvoiceSettlement, + check_invoice_payment, +) + + +def _lightning_invoice(**overrides: object) -> LightningInvoice: + suffix = uuid.uuid4().hex + values = { + "id": f"invoice-{suffix}", + "bolt11": f"lnbc-{suffix}", + "amount_sats": 100, + "description": "settlement test", + "payment_hash": f"quote-{suffix}", + "status": "pending", + "purpose": "create", + "mint_url": "http://mint:3338", + "expires_at": int(time.time()) + 3600, + } + values.update(overrides) + return LightningInvoice(**values) # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_invoice_read_transaction_closes_before_external_mint_io( + integration_session: AsyncSession, +) -> None: + invoice = _lightning_invoice() + integration_session.add(invoice) + await integration_session.commit() + stored = await integration_session.get(LightningInvoice, invoice.id) + assert stored is not None + + wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=False))) + + async def get_wallet_without_open_db_transaction( + *args: object, **kwargs: object + ) -> Mock: + assert not integration_session.in_transaction() + return wallet + + with patch( + "routstr.lightning.get_wallet", side_effect=get_wallet_without_open_db_transaction + ): + await check_invoice_payment(stored, integration_session) + + assert not integration_session.in_transaction() + + +@pytest.mark.asyncio +async def test_separate_sessions_cas_topup_credit_exactly_once( + integration_engine: object, +) -> None: + key_hash = uuid.uuid4().hex + invoice = _lightning_invoice( + purpose="topup", + api_key_hash=key_hash, + amount_sats=100, + ) + key = ApiKey( + hashed_key=key_hash, + balance=100_000, + refund_currency="sat", + refund_mint_url="http://mint:3338", + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as seed: + seed.add(key) + seed.add(invoice) + await seed.commit() + + snapshot_a = _InvoiceSettlement.from_invoice(invoice) + snapshot_b = _InvoiceSettlement.from_invoice(invoice) + async with ( + AsyncSession(integration_engine, expire_on_commit=False) as session_a, + AsyncSession(integration_engine, expire_on_commit=False) as session_b, + ): + results = await asyncio.gather( + _finalize_invoice_settlement(snapshot_a, session_a, 1_700_000_000), + _finalize_invoice_settlement(snapshot_b, session_b, 1_700_000_001), + ) + + assert sorted(settled for settled, _ in results) == [False, True] + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + stored_invoice = await verify.get(LightningInvoice, invoice.id) + stored_key = await verify.get(ApiKey, key_hash) + assert stored_invoice is not None + assert stored_invoice.status == "paid" + assert stored_key is not None + assert stored_key.balance == 200_000 + + +@pytest.mark.asyncio +async def test_topup_atomic_increment_preserves_concurrent_balance_mutation( + integration_engine: object, +) -> None: + key_hash = uuid.uuid4().hex + invoice = _lightning_invoice( + purpose="topup", api_key_hash=key_hash, amount_sats=100 + ) + key = ApiKey( + hashed_key=key_hash, + balance=100_000, + refund_currency="sat", + refund_mint_url="http://mint:3338", + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as seed: + seed.add(key) + seed.add(invoice) + await seed.commit() + + async def debit_balance(session: AsyncSession) -> None: + result = await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == key_hash) + .values(balance=col(ApiKey.balance) - 10_000) + .execution_options(synchronize_session=False) + ) + assert result.rowcount == 1 + await session.commit() + + snapshot = _InvoiceSettlement.from_invoice(invoice) + async with ( + AsyncSession(integration_engine, expire_on_commit=False) as settlement, + AsyncSession(integration_engine, expire_on_commit=False) as debit, + ): + settlement_result, _ = await asyncio.gather( + _finalize_invoice_settlement(snapshot, settlement, 1_700_000_000), + debit_balance(debit), + ) + + assert settlement_result[0] + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + stored_key = await verify.get(ApiKey, key_hash) + assert stored_key is not None + assert stored_key.balance == 190_000 + + +@pytest.mark.asyncio +async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry( + integration_engine: object, +) -> None: + key_hash = uuid.uuid4().hex + invoice = _lightning_invoice( + purpose="topup", + api_key_hash=key_hash, + amount_sats=100, + ) + key = ApiKey( + hashed_key=key_hash, + balance=100_000, + refund_currency="sat", + refund_mint_url="http://mint:3338", + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as seed: + seed.add(key) + seed.add(invoice) + await seed.commit() + + snapshot = _InvoiceSettlement.from_invoice(invoice) + async with AsyncSession(integration_engine, expire_on_commit=False) as failed: + with patch.object( + failed, "commit", AsyncMock(side_effect=Exception("db unavailable")) + ): + with pytest.raises(Exception, match="db unavailable"): + await _finalize_invoice_settlement(snapshot, failed, 1_700_000_000) + + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + pending = await verify.get(LightningInvoice, invoice.id) + unchanged = await verify.get(ApiKey, key_hash) + assert pending is not None + assert pending.status == "pending" + assert unchanged is not None + assert unchanged.balance == 100_000 + + async with AsyncSession(integration_engine, expire_on_commit=False) as retry: + settled, _ = await _finalize_invoice_settlement( + snapshot, retry, 1_700_000_001 + ) + assert settled + + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + paid = await verify.get(LightningInvoice, invoice.id) + credited = await verify.get(ApiKey, key_hash) + assert paid is not None + assert paid.status == "paid" + assert credited is not None + assert credited.balance == 200_000 + + +@pytest.mark.asyncio +async def test_check_invoice_payment_retries_after_mint_success_and_db_failure( + integration_engine: object, +) -> None: + key_hash = uuid.uuid4().hex + invoice = _lightning_invoice( + purpose="topup", api_key_hash=key_hash, amount_sats=100 + ) + key = ApiKey( + hashed_key=key_hash, + balance=100_000, + refund_currency="sat", + refund_mint_url="http://mint:3338", + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as seed: + seed.add(key) + seed.add(invoice) + await seed.commit() + + wallet = Mock( + proofs=[], + keysets={"keyset-1": Mock()}, + load_proofs=AsyncMock(), + get_mint_quote=AsyncMock(return_value=Mock(paid=True)), + restore_tokens_for_keyset=AsyncMock(), + ) + + async def mint(amount: int, quote_id: str) -> list[Proof]: + proofs = [Proof(amount=amount, mint_id=quote_id)] + wallet.proofs.extend(proofs) + return proofs + + wallet.mint = AsyncMock(side_effect=mint) + + async with AsyncSession(integration_engine, expire_on_commit=False) as failed: + stored = await failed.get(LightningInvoice, invoice.id) + assert stored is not None + real_commit = failed.commit + commit_count = 0 + + async def fail_final_commit() -> None: + nonlocal commit_count + commit_count += 1 + if commit_count == 2: + raise Exception("db unavailable") + await real_commit() + + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch.object(failed, "commit", AsyncMock(side_effect=fail_final_commit)), + ): + await check_invoice_payment(stored, failed) + + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + pending = await verify.get(LightningInvoice, invoice.id) + unchanged = await verify.get(ApiKey, key_hash) + assert pending is not None + assert pending.status == "pending" + assert unchanged is not None + assert unchanged.balance == 100_000 + + async with AsyncSession(integration_engine, expire_on_commit=False) as retry: + stored = await retry.get(LightningInvoice, invoice.id) + assert stored is not None + with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)): + await check_invoice_payment(stored, retry) + + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + paid = await verify.get(LightningInvoice, invoice.id) + credited = await verify.get(ApiKey, key_hash) + assert paid is not None + assert paid.status == "paid" + assert credited is not None + assert credited.balance == 200_000 + + wallet.mint.assert_awaited_once_with(100, quote_id=invoice.payment_hash) + wallet.restore_tokens_for_keyset.assert_not_awaited() diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py new file mode 100644 index 00000000..22590a86 --- /dev/null +++ b/tests/unit/test_lightning_settlement.py @@ -0,0 +1,206 @@ +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock, patch + +import httpx +import pytest +from cashu.core.base import Proof + +from routstr.lightning import ( + _invoice_settlement_locks, + _is_outputs_already_signed, + _mint_invoice_quote, + check_invoice_payment, +) +from routstr.wallet import Wallet + + +def _invoice(**overrides: object) -> SimpleNamespace: + values = { + "id": "invoice-1", + "payment_hash": "quote-1", + "amount_sats": 100, + "purpose": "create", + "status": "pending", + "paid_at": None, + "api_key_hash": None, + "mint_url": "http://mint:3338", + "balance_limit": None, + "balance_limit_reset": None, + "validity_date": None, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def _proof(amount: int, mint_id: str, *, reserved: bool = False) -> Proof: + return Proof(amount=amount, mint_id=mint_id, reserved=reserved) + + +def _recovery_wallet( + error: Exception, + *, + proofs_before: list[Proof] | None = None, + proofs_after: list[Proof] | None = None, +) -> Mock: + async def load_proofs(*, reload: bool) -> None: + if wallet.load_proofs.await_count >= 2 and proofs_after is not None: + wallet.proofs = list(proofs_after) + + wallet = Mock( + mint=AsyncMock(side_effect=error), + keysets={"keyset-1": Mock()}, + restore_tokens_for_keyset=AsyncMock(), + load_proofs=AsyncMock(side_effect=load_proofs), + proofs=list(proofs_before or []), + ) + return wallet + + +@pytest.mark.asyncio +async def test_invoice_mint_recovers_quote_linked_outputs_already_signed() -> None: + invoice = _invoice() + wallet = _recovery_wallet( + Exception("Mint Error: outputs have already been signed before (Code: 11003)"), + proofs_after=[_proof(100, "quote-1")], + ) + + await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type] + + wallet.restore_tokens_for_keyset.assert_awaited_once_with( + "keyset-1", to=1, batch=25 + ) + assert wallet.load_proofs.await_count == 2 + + +@pytest.mark.asyncio +async def test_invoice_mint_accepts_preloaded_quote_linked_proofs() -> None: + invoice = _invoice() + wallet = _recovery_wallet( + Exception("must not mint"), + proofs_before=[_proof(64, "quote-1"), _proof(36, "quote-1")], + ) + + await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type] + + wallet.mint.assert_not_awaited() + wallet.restore_tokens_for_keyset.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_invoice_mint_does_not_accept_unrelated_11003_text() -> None: + invoice = _invoice() + error = Exception("backend request 11003 failed") + wallet = _recovery_wallet(error) + + with pytest.raises(Exception) as caught: + await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type] + + assert caught.value is error + wallet.restore_tokens_for_keyset.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_installed_cashu_error_shape_recognizes_realistic_11003_phrase() -> None: + request = httpx.Request("POST", "http://mint:3338/v1/mint/bolt11") + response = httpx.Response( + 400, + request=request, + json={"detail": "outputs have already been signed before", "code": 11003}, + ) + + with pytest.raises(Exception) as caught: + Wallet.raise_on_error_request(response) + + assert _is_outputs_already_signed(caught.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("recovered", [0, 99]) +async def test_invoice_mint_rejects_empty_or_short_quote_recovery( + recovered: int, +) -> None: + invoice = _invoice() + wallet = _recovery_wallet( + Exception("Mint Error: outputs already signed (Code: 11003)"), + proofs_after=[_proof(recovered, "quote-1")] if recovered else [], + ) + + with pytest.raises(RuntimeError, match="expected at least 100"): + await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_invoice_mint_rejects_unrelated_concurrent_balance_growth() -> None: + invoice = _invoice() + wallet = _recovery_wallet( + Exception("Mint Error: outputs already signed (Code: 11003)"), + proofs_after=[_proof(10_000, "different-quote")], + ) + + with pytest.raises(RuntimeError, match="quote-linked recovery returned 0"): + await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_non_pending_invoice_is_not_minted() -> None: + _invoice_settlement_locks.clear() + invoice = _invoice(status="expired") + session = AsyncMock() + + with patch("routstr.lightning.get_wallet", AsyncMock()) as get_wallet: + await check_invoice_payment(invoice, session) # type: ignore[arg-type] + + get_wallet.assert_not_awaited() + session.commit.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_ambiguous_invoice_mint_timeout_does_not_expose_paid() -> None: + _invoice_settlement_locks.clear() + invoice = _invoice() + session = AsyncMock() + wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True))) + + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch( + "routstr.lightning._mint_invoice_quote", + AsyncMock(side_effect=httpx.TimeoutException("response lost")), + ), + patch("routstr.lightning._reload_invoice_view", AsyncMock()), + ): + await check_invoice_payment(invoice, session) # type: ignore[arg-type] + + assert invoice.status == "pending" + session.rollback.assert_awaited_once() + # One commit closes the initial read transaction before external I/O. + session.commit.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_concurrent_invoice_checks_finalize_once_in_process() -> None: + _invoice_settlement_locks.clear() + invoice = _invoice() + session = AsyncMock() + wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True))) + + async def refresh(obj: SimpleNamespace) -> None: + return None + + session.refresh = AsyncMock(side_effect=refresh) + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning._mint_invoice_quote", AsyncMock()), + patch( + "routstr.lightning._finalize_invoice_settlement", + AsyncMock(return_value=(True, "b" * 64)), + ) as finalize, + ): + await asyncio.gather( + check_invoice_payment(invoice, session), # type: ignore[arg-type] + check_invoice_payment(invoice, session), # type: ignore[arg-type] + ) + + assert invoice.status == "paid" + finalize.assert_awaited_once() diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index a91991f9..fe0cd573 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -13,8 +13,10 @@ from routstr.payment.helpers import ( # noqa: E402 ) -def test_mint_fee_allowance_reduces_admission_cost_by_ten_percent() -> None: - assert apply_mint_fee_allowance(124_886) == 112_398 +def test_mint_fee_allowance_reserves_five_percent_fallback_headroom() -> None: + # Interim policy: Routstr may pay hidden cross-mint Lightning fees when a + # trusted-mint fallback is required. + assert apply_mint_fee_allowance(124_886) == 118_642 def test_mint_fee_allowance_never_drops_below_minimum() -> None: diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 886331c3..512dfa31 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -422,4 +422,4 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: with pytest.raises(asyncio.CancelledError): await proxy_module.proxy(request, "v1/chat/completions", session=session) - revert_mock.assert_awaited_once_with(key, session, 900, reservation_snapshot) + revert_mock.assert_awaited_once_with(key, session, 950, reservation_snapshot) diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index d64d95e3..65b50a73 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -406,4 +406,4 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: assert RAW_ORG_ID not in serialized assert "org-[REDACTED]" in serialized # Single upstream failed -> reservation reverted exactly once (no double-charge). - revert_mock.assert_awaited_once_with(key, session, 900, reservation) + revert_mock.assert_awaited_once_with(key, session, 950, reservation) diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 49369f14..257b6132 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1726,20 +1726,48 @@ async def test_mint_operation_honors_retry_after_as_minimum() -> None: @pytest.mark.asyncio -async def test_mint_operation_timeout_includes_adaptive_cooldown() -> None: +async def test_mint_operation_timeout_excludes_adaptive_cooldown() -> None: from routstr.core.settings import settings from routstr.wallet import _mint_operation, _MintRateGuard - operation = AsyncMock(return_value="unexpected") - with patch.object(settings, "mint_max_concurrency", 1): + operation = AsyncMock(return_value="ok") + with ( + patch.object(settings, "mint_max_concurrency", 1), + patch.object(settings, "mint_operation_timeout_seconds", 0.01), + patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + ): guard = _MintRateGuard.get("http://mint:3338") - assert guard is not None guard.apply_cooldown(60) - with patch.object(settings, "mint_operation_timeout_seconds", 0.01): - with pytest.raises(httpx.TimeoutException, match="total timeout"): - await _mint_operation(operation, mint_url="http://mint:3338") + assert await _mint_operation(operation, mint_url="http://mint:3338") == "ok" - operation.assert_not_awaited() + sleep.assert_awaited_once() + operation.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_default_timeout_allows_retry_after_rate_limit_cooldown() -> None: + from routstr.core.settings import settings + from routstr.wallet import _mint_operation + + request = httpx.Request("POST", "http://mint:3338/v1/mint/quote/bolt11") + response = httpx.Response(429, request=request) + operation = AsyncMock( + side_effect=[ + httpx.HTTPStatusError( + "rate limited", request=request, response=response + ), + "ok", + ] + ) + with ( + patch.object(settings, "mint_retry_max_attempts", 3), + patch.object(settings, "mint_operation_timeout_seconds", 30), + patch.object(settings, "mint_max_concurrency", 1), + patch("routstr.wallet.asyncio.sleep", AsyncMock()), + ): + assert await _mint_operation(operation, mint_url="http://mint:3338") == "ok" + + assert operation.await_count == 2 @pytest.mark.asyncio @@ -1984,27 +2012,26 @@ async def test_swap_falls_back_when_primary_wallet_cannot_load() -> None: @pytest.mark.asyncio -async def test_lightning_mint_fallback_on_429() -> None: - """A 429 from the primary mint should trigger fallback to a secondary, - not just transport errors.""" +async def test_lightning_mint_fallback_on_cashu_json_429() -> None: + """The real Cashu JSON-error adapter preserves 429 for fallback.""" from routstr.core.settings import settings from routstr.lightning import _request_mint_with_fallback + from routstr.wallet import MintRateLimitedError, Wallet primary = "http://primary:3338" secondary = "http://secondary:3338" - mock_resp = Mock(status_code=429, headers={}) - mock_resp.raise_for_status = Mock( - side_effect=httpx.HTTPStatusError( - "rate limited", request=Mock(), response=mock_resp - ) + request = httpx.Request("POST", f"{primary}/v1/mint/quote/bolt11") + response = httpx.Response( + 429, + request=request, + json={"detail": "too many requests", "code": 0}, ) + with pytest.raises(MintRateLimitedError) as captured: + Wallet.raise_on_error_request(response) + mock_primary_wallet = Mock() - mock_primary_wallet.request_mint = AsyncMock( - side_effect=httpx.HTTPStatusError( - "rate limited", request=Mock(), response=mock_resp - ) - ) + mock_primary_wallet.request_mint = AsyncMock(side_effect=captured.value) mock_quote = Mock(request="lnbc1secondary", quote="quote_secondary") mock_secondary_wallet = Mock() From 48c11eb7bceb1edda42373667211106f2afa02a1 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 27 Jul 2026 00:07:36 +0200 Subject: [PATCH 35/46] fix Lightning settlement test typing --- tests/integration/test_lightning_settlement.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tests/integration/test_lightning_settlement.py b/tests/integration/test_lightning_settlement.py index dff37c48..ff9e80b2 100644 --- a/tests/integration/test_lightning_settlement.py +++ b/tests/integration/test_lightning_settlement.py @@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, Mock, patch import pytest from cashu.core.base import Proof +from sqlalchemy.ext.asyncio import AsyncEngine from sqlmodel import col, update from sqlmodel.ext.asyncio.session import AsyncSession @@ -61,7 +62,7 @@ async def test_invoice_read_transaction_closes_before_external_mint_io( @pytest.mark.asyncio async def test_separate_sessions_cas_topup_credit_exactly_once( - integration_engine: object, + integration_engine: AsyncEngine, ) -> None: key_hash = uuid.uuid4().hex invoice = _lightning_invoice( @@ -103,7 +104,7 @@ async def test_separate_sessions_cas_topup_credit_exactly_once( @pytest.mark.asyncio async def test_topup_atomic_increment_preserves_concurrent_balance_mutation( - integration_engine: object, + integration_engine: AsyncEngine, ) -> None: key_hash = uuid.uuid4().hex invoice = _lightning_invoice( @@ -149,7 +150,7 @@ async def test_topup_atomic_increment_preserves_concurrent_balance_mutation( @pytest.mark.asyncio async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry( - integration_engine: object, + integration_engine: AsyncEngine, ) -> None: key_hash = uuid.uuid4().hex invoice = _lightning_invoice( @@ -201,7 +202,7 @@ async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry( @pytest.mark.asyncio async def test_check_invoice_payment_retries_after_mint_success_and_db_failure( - integration_engine: object, + integration_engine: AsyncEngine, ) -> None: key_hash = uuid.uuid4().hex invoice = _lightning_invoice( From e2f89a26454debb3b44860508209457eaff50ab6 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 27 Jul 2026 23:17:13 +0200 Subject: [PATCH 36/46] 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, From c75dee147ab6378bdf2f3c1e22368d668bcfd951 Mon Sep 17 00:00:00 2001 From: thefux Date: Tue, 28 Jul 2026 00:08:47 +0000 Subject: [PATCH 37/46] fix: set routstr-core port to 8011 to avoid Portainer conflict on 8000 --- compose.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/compose.yml b/compose.yml index 86ee26b7..2d2496b5 100644 --- a/compose.yml +++ b/compose.yml @@ -29,7 +29,7 @@ services: environment: - TOR_PROXY_URL=socks5://tor:9050 ports: - - 8000:8000 + - 8011:8000 extra_hosts: # Needed to access locally running models - "host.docker.internal:host-gateway" restart: unless-stopped From bb2a05b67ce0bb21c5b055a56041be9ea53ea6e0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 29 Jul 2026 00:08:34 +0200 Subject: [PATCH 38/46] add url and specific model infos to paht --- docs/api/endpoints.md | 17 ++-- routstr/core/db.py | 5 +- routstr/upstream/model_paths.py | 65 ++++++++++--- tests/unit/test_model_paths.py | 167 +++++++++++++++++++++++++------- 4 files changed, 195 insertions(+), 59 deletions(-) diff --git a/docs/api/endpoints.md b/docs/api/endpoints.md index ca54b8e0..23b019c7 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=anthropic-primary", + "path": "url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=anthropic%2Fclaude-sonnet-4", "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, "endpoint": null }, { - "path": "provider=openrouter-main&endpoint=google-vertex%2Fus", + "path": "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4&endpoint=google-vertex%2Fus", "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, "endpoint": {"tag": "google-vertex/us", "name": "Google"} } @@ -362,12 +362,13 @@ 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 -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 -not silently fall back to another backend. +unchanged rather than parsing or reconstructing it. It identifies the exact +configured route with `url`, `provider-id`, and `model-id`. To avoid exposing +private network details, a configured private IP address or any URL with an +explicit port is advertised as `http://localhost`. OpenRouter routes additionally +preserve the exact machine-readable endpoint `tag`. Provider slugs/types and +endpoint names remain display data. When request-side selection is implemented, +an endpoint tag must not silently fall back to another backend. ### List Paths for One Model diff --git a/routstr/core/db.py b/routstr/core/db.py index 3edfe140..e4794d4a 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -338,7 +338,10 @@ 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 slug and optional endpoint tag" + description=( + "Opaque selector containing upstream URL, provider ID, model ID, " + "and optional endpoint tag" + ) ) provider_slug: str = Field( description="Public slug of the configured upstream provider" diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 40022aca..020f0fbb 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -5,22 +5,22 @@ 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 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:: +upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter +endpoint, its machine-readable tag:: - provider=anthropic-primary - provider=openrouter-main&endpoint=google-vertex%2Fus + url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=claude-sonnet-4 + url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=claude-sonnet-4&endpoint=google-vertex%2Fus """ from __future__ import annotations import asyncio +import ipaddress import random import time from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Callable -from urllib.parse import urlencode +from urllib.parse import urlencode, urlsplit import httpx from sqlalchemy.dialects.sqlite import insert @@ -66,11 +66,12 @@ class EndpointIdentity: @dataclass(frozen=True) class ConfiguredProviderIdentity: - """Public-safe identity of one configured upstream provider.""" + """Public identity of one configured upstream provider.""" id: int slug: str provider_type: str + base_url: str @dataclass(frozen=True) @@ -92,9 +93,38 @@ class ProviderPathSnapshot: preserve_model_ids: frozenset[str] = frozenset() -def encode_model_path(provider_slug: str, endpoint_tag: str | None = None) -> str: - """Encode a stable opaque selector without exposing upstream URLs.""" - components = [("provider", provider_slug)] +def public_provider_url(base_url: str) -> str: + """Mask private IP addresses and URLs with explicit ports.""" + parsed = urlsplit(base_url) + try: + if parsed.port is not None: + return "http://localhost" + except ValueError: + # An invalid explicit port must not accidentally leak through. + return "http://localhost" + + hostname = parsed.hostname + if hostname is None: + return base_url + try: + address = ipaddress.ip_address(hostname) + except ValueError: + return base_url + return "http://localhost" if address.is_private else base_url + + +def encode_model_path( + base_url: str, + provider_id: int, + model_id: str, + endpoint_tag: str | None = None, +) -> str: + """Encode the complete upstream route selector advertised to clients.""" + components: list[tuple[str, str | int]] = [ + ("url", base_url), + ("provider-id", provider_id), + ("model-id", model_id), + ] if endpoint_tag: components.append(("endpoint", endpoint_tag)) return urlencode(components) @@ -297,6 +327,7 @@ async def _load_model_visibility() -> tuple[ id=provider.id, slug=provider.slug or f"provider-{provider.id}", provider_type=provider.provider_type, + base_url=public_provider_url(provider.base_url), ) for model in provider.models: key = (model.id.lower(), provider.id) @@ -374,9 +405,12 @@ async def _collect_provider_paths( models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) def _base_path(model: object) -> DiscoveredPath: + model_id = exposed_model_id(model) return DiscoveredPath( - model_id=exposed_model_id(model), - path=encode_model_path(provider_identity.slug), + model_id=model_id, + path=encode_model_path( + provider_identity.base_url, provider_identity.id, model_id + ), provider=provider_identity, ) @@ -410,7 +444,12 @@ async def _collect_provider_paths( paths.extend( DiscoveredPath( model_id=model_id, - path=encode_model_path(provider_identity.slug, endpoint.tag), + path=encode_model_path( + provider_identity.base_url, + provider_identity.id, + model_id, + endpoint.tag, + ), provider=provider_identity, endpoint_tag=endpoint.tag, endpoint_name=endpoint.provider_name, diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 70986dfd..1bbe25e0 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -213,8 +213,19 @@ def _ids_of(payload: dict) -> set[str]: return {entry["id"] for entry in payload["data"]} +def _expected_path( + provider_id: int, + model_id: str, + endpoint_tag: str | None = None, +) -> str: + return mp.encode_model_path( + f"https://provider-{provider_id}", provider_id, model_id, endpoint_tag + ) + + def _path_entry( provider_id: int, + model_id: str, *, provider_slug: str | None = None, provider_type: str | None = None, @@ -225,7 +236,7 @@ def _path_entry( if endpoint_tag or endpoint_name: endpoint = {"tag": endpoint_tag, "name": endpoint_name} return { - "path": mp.encode_model_path(provider_slug or f"p{provider_id}", endpoint_tag), + "path": _expected_path(provider_id, model_id, endpoint_tag), "provider": { "id": provider_id, "slug": provider_slug or f"p{provider_id}", @@ -253,10 +264,38 @@ 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_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" +def test_public_provider_url_masks_private_addresses_and_explicit_ports() -> None: + assert mp.public_provider_url("http://192.168.1.10/v1") == "http://localhost" + assert mp.public_provider_url("http://10.0.0.5:11434/v1") == "http://localhost" + assert mp.public_provider_url("http://[fd00::1]/v1") == "http://localhost" + assert mp.public_provider_url("https://api.example.com:8443/v1") == ( + "http://localhost" + ) + + +def test_public_provider_url_preserves_public_urls_without_ports() -> None: + assert mp.public_provider_url("https://openrouter.ai/api/v1") == ( + "https://openrouter.ai/api/v1" + ) + assert mp.public_provider_url("http://localhost") == "http://localhost" + + +def test_encode_model_path_includes_complete_route_identity() -> None: + assert mp.encode_model_path( + "https://openrouter.ai/api/v1", 42, "anthropic/claude-sonnet-4" + ) == ( + "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1" + "&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4" + ) + assert mp.encode_model_path( + "https://openrouter.ai/api/v1", + 42, + "anthropic/claude-sonnet-4", + "google-vertex/us-east5", + ) == ( + "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1" + "&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4" + "&endpoint=google-vertex%2Fus-east5" ) @@ -322,10 +361,36 @@ async def test_direct_provider_single_path_uses_provider_type( ) await mp.refresh_model_paths([provider]) payload = await mp.get_all_model_paths() - assert payload["data"] == [{"id": "claude-opus-4.6", "paths": [_path_entry(1)]}] + assert payload["data"] == [ + {"id": "claude-opus-4.6", "paths": [_path_entry(1, "claude-opus-4.6")]} + ] assert payload["updated_at"] is not None +@pytest.mark.asyncio +async def test_direct_path_masks_private_configured_provider_url( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.base_url = "http://192.168.1.10:11434/v1" + session.add(provider_row) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="http://192.168.1.10:11434/v1", + models=[_model("local-model")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "local-model") == { + mp.encode_model_path("http://localhost", 1, "local-model") + } + + @pytest.mark.asyncio async def test_direct_path_stores_exposed_model_id( patched_session: AsyncEngine, @@ -397,7 +462,9 @@ 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("p1")} + assert _paths_of(payload, "shared-model") == { + _expected_path(1, "shared-model") + } @pytest.mark.asyncio @@ -431,8 +498,12 @@ 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("p1")} - assert _paths_of(payload, "private-alias") == {mp.encode_model_path("p2")} + assert _paths_of(payload, "shared-model") == { + _expected_path(1, "shared-model") + } + assert _paths_of(payload, "private-alias") == { + _expected_path(2, "private-alias") + } @pytest.mark.asyncio @@ -500,7 +571,7 @@ async def test_refresh_model_paths_uses_db_forwarded_alias( await mp.refresh_model_paths([provider]) assert (await mp.get_all_model_paths())["data"] == [ - {"id": "public-alias", "paths": [_path_entry(1)]} + {"id": "public-alias", "paths": [_path_entry(1, "public-alias")]} ] @@ -522,7 +593,10 @@ async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cac await mp.refresh_model_paths([provider]) assert (await mp.get_all_model_paths())["data"] == [ - {"id": "public-deployment", "paths": [_path_entry(1)]} + { + "id": "public-deployment", + "paths": [_path_entry(1, "public-deployment")], + } ] @@ -684,9 +758,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("p2"), - mp.encode_model_path("p2", "google-vertex/eu"), - mp.encode_model_path("p2", "google-vertex/us"), + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "google-vertex/eu"), + _expected_path(2, "claude-opus-4.6", "google-vertex/us"), } assert { item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] @@ -712,8 +786,8 @@ async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path("p2"), - mp.encode_model_path("p2", "openrouter"), + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "openrouter"), } @@ -734,8 +808,8 @@ async def test_generic_provider_with_openrouter_base_url_discovers( paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") assert paths == { - mp.encode_model_path("p1"), - mp.encode_model_path("p1", "anthropic"), + _expected_path(1, "claude-opus-4.6"), + _expected_path(1, "claude-opus-4.6", "anthropic"), } @@ -792,9 +866,9 @@ async def test_partial_failure_upserts_collapsed_public_model_paths( 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"), + _expected_path(2, "shared"), + _expected_path(2, "shared", "anthropic"), + _expected_path(2, "shared", "google"), } @@ -811,7 +885,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("p2", "anthropic") in before + assert _expected_path(2, "claude-opus-4.6", "anthropic") in before def _network_down(request: httpx.Request) -> httpx.Response: raise httpx.ConnectError("network down", request=request) @@ -835,8 +909,8 @@ 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("p2"), - mp.encode_model_path("p2", "anthropic"), + _expected_path(2, "m0"), + _expected_path(2, "m0", "anthropic"), } assert _paths_of(await mp.get_all_model_paths(), "m0") == expected @@ -906,10 +980,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("p2"), - mp.encode_model_path("p2", "anthropic"), - mp.encode_model_path("p4"), - mp.encode_model_path("p4", "anthropic"), + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "anthropic"), + _expected_path(4, "claude-opus-4.6"), + _expected_path(4, "claude-opus-4.6", "anthropic"), } @@ -972,8 +1046,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("p1"), - mp.encode_model_path("p2"), + _expected_path(1, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6"), } assert "canonical_id" not in entry assert all("canonical_id" not in p for p in entry["paths"]) @@ -1000,7 +1074,10 @@ async def test_get_all_model_paths_keeps_distinct_configured_providers( assert (await mp.get_all_model_paths())["data"] == [ { "id": "claude-opus-4.6", - "paths": [_path_entry(1), _path_entry(2)], + "paths": [ + _path_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ], } ] @@ -1032,7 +1109,10 @@ async def test_get_paths_for_model_returns_route_identity( await _seed_two_provider_shared_model(patched_session) payload = await mp.get_paths_for_model("claude-opus-4.6") - assert payload["data"] == [_path_entry(1), _path_entry(2)] + assert payload["data"] == [ + _path_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ] assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] @@ -1048,7 +1128,9 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( ) await mp.refresh_model_paths([provider]) - assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [_path_entry(4)] + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ + _path_entry(4, "glm-5v-turbo") + ] @pytest.mark.asyncio @@ -1072,7 +1154,10 @@ async def test_get_paths_for_model_requires_exact_advertised_id( short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] - assert short_paths == [_path_entry(4), _path_entry(7)] + assert short_paths == [ + _path_entry(4, "deepseek-v4-pro"), + _path_entry(7, "deepseek-v4-pro"), + ] assert prefixed_paths == [] @@ -1092,7 +1177,7 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ - _path_entry(1) + _path_entry(1, "fireworks/models/glm-5") ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" @@ -1284,9 +1369,10 @@ def test_model_paths_endpoint_returns_all_paths( { "id": "claude-opus-4.6", "paths": [ - _path_entry(1), + _path_entry(1, "claude-opus-4.6"), _path_entry( 2, + "claude-opus-4.6", endpoint_tag="google-vertex/us", endpoint_name="Google", ), @@ -1369,7 +1455,14 @@ def test_model_paths_for_model_endpoint_accepts_slash_model_id( calls: list[str] = [] expected = { - "data": [_path_entry(2, endpoint_tag="anthropic", endpoint_name="Anthropic")], + "data": [ + _path_entry( + 2, + "anthropic/claude-opus-4.6", + endpoint_tag="anthropic", + endpoint_name="Anthropic", + ) + ], "updated_at": None, } From f8adaee362b736aae3229b32c903c1a76fde712f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 30 Jul 2026 01:14:36 +0200 Subject: [PATCH 39/46] revert: restore default compose port --- compose.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/compose.yml b/compose.yml index 2d2496b5..86ee26b7 100644 --- a/compose.yml +++ b/compose.yml @@ -29,7 +29,7 @@ services: environment: - TOR_PROXY_URL=socks5://tor:9050 ports: - - 8011:8000 + - 8000:8000 extra_hosts: # Needed to access locally running models - "host.docker.internal:host-gateway" restart: unless-stopped From 3befe063f4e0b3efc34ed8b0de10f08d7696e358 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 30 Jul 2026 01:19:35 +0200 Subject: [PATCH 40/46] fix: annotate lightning settlement test session --- tests/unit/test_lightning_settlement.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py index 7f748e9e..c44954f2 100644 --- a/tests/unit/test_lightning_settlement.py +++ b/tests/unit/test_lightning_settlement.py @@ -1,4 +1,5 @@ import asyncio +from collections.abc import AsyncIterator from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch @@ -192,7 +193,7 @@ async def test_concurrent_invoice_checks_finalize_once_in_process() -> None: session.refresh = AsyncMock(side_effect=refresh) @asynccontextmanager - async def owned_session(): + async def owned_session() -> AsyncIterator[AsyncMock]: yield AsyncMock() with ( From 2ec6b2720092511ffcd80a9b774612076571a965 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 31 Jul 2026 02:10:00 +0200 Subject: [PATCH 41/46] fix: resolve mint fallback review comments --- .env.example | 3 + docs/provider/configuration.md | 9 + routstr/balance.py | 4 +- routstr/lightning.py | 52 +- routstr/mint.py | 338 ++++++++++ routstr/payment/helpers.py | 12 - routstr/payment/lnurl.py | 66 +- routstr/proxy.py | 5 - routstr/wallet.py | 613 ++++++------------ .../integration/test_insufficient_balance.py | 31 +- tests/integration/test_swap_fee_retry.py | 7 +- tests/unit/test_fetch_all_balances.py | 6 +- tests/unit/test_lightning_settlement.py | 2 + tests/unit/test_lnurl_melt_timeout.py | 144 ++-- tests/unit/test_melt_reconciliation.py | 91 +++ tests/unit/test_mint.py | 65 ++ tests/unit/test_payment_helpers.py | 16 +- tests/unit/test_stale_reservations.py | 2 +- tests/unit/test_upstream_rate_limit.py | 2 +- tests/unit/test_wallet.py | 101 ++- 20 files changed, 957 insertions(+), 612 deletions(-) create mode 100644 routstr/mint.py create mode 100644 tests/unit/test_melt_reconciliation.py create mode 100644 tests/unit/test_mint.py diff --git a/.env.example b/.env.example index 9f0c0bbc..3d12ee5c 100644 --- a/.env.example +++ b/.env.example @@ -45,6 +45,9 @@ ROUTSTR_SECRET_KEY= # ENABLE_ANALYTICS_SHARING=true # CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me" # MINT_OPERATION_CONCURRENCY=4 +# MINT_OPERATION_TIMEOUT_SECONDS=30 +# MINT_MAX_CONCURRENCY=4 +# MINT_RETRY_MAX_ATTEMPTS=3 # RECEIVE_LN_ADDRESS= # REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS=900 diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 33efabb8..eee669c4 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -136,6 +136,10 @@ Use environment variables for: | `NSEC` | Legacy seed for the Nostr private key (otherwise set from the admin UI) | — | | `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` | | `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` | +| `MINT_OPERATION_CONCURRENCY` | Concurrent mint/unit balance reads | `4` | +| `MINT_OPERATION_TIMEOUT_SECONDS` | Per-attempt timeout for mint network calls | `30` | +| `MINT_MAX_CONCURRENCY` | Concurrent operations allowed per mint (`0` disables the limit) | `4` | +| `MINT_RETRY_MAX_ATTEMPTS` | Retries after a timeout or HTTP 429 (`0` disables retries) | `3` | | `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — | | `MIN_PAYOUT_SAT` | Min payout balance in sats (applies to all mints) | `210` | | `PAYOUT_INTERVAL_SECONDS` | Payout loop interval (seconds) | `900` | @@ -143,6 +147,11 @@ Use environment variables for: | `CORS_ORIGINS` | Allowed CORS origins | `*` | | `RELAYS` | Nostr relays (comma-separated) | (default set) | +Mint HTTP 429 responses create a per-mint cooldown. Operations that already hold +Routstr's wallet mutation lock fail fast during that cooldown instead of waiting +while blocking every other wallet mutation. Callers receive an error and may retry +later; the current response does not include the cooldown duration. + ### Priority Environment variables are read on startup. Dashboard settings override them and persist in the database. Once you change a setting in the dashboard, the env var is ignored for that setting. diff --git a/routstr/balance.py b/routstr/balance.py index 00894608..b106aebd 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -215,7 +215,7 @@ async def topup_wallet_endpoint( raise HTTPException(status_code=400, detail="Invalid token format") source_mint = token_mint_url(cashu_token, "unknown") - logger.warning( + logger.info( "Cashu wallet top-up started", extra={ "event": "cashu_topup_started", @@ -259,7 +259,7 @@ async def topup_wallet_endpoint( ) raise HTTPException(status_code=status_code, detail=message) - logger.warning( + logger.info( "Cashu wallet top-up completed", extra={ "event": "cashu_topup_completed", diff --git a/routstr/lightning.py b/routstr/lightning.py index 6fa4976e..41dc775f 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -3,8 +3,9 @@ import hashlib import re import secrets import time +from contextlib import asynccontextmanager from dataclasses import dataclass -from typing import Any +from typing import Any, AsyncGenerator from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel, Field @@ -15,11 +16,13 @@ from sqlmodel.ext.asyncio.session import AsyncSession from .core.db import ApiKey, LightningInvoice, create_session, get_session from .core.logging import get_logger from .core.settings import settings +from .mint import ( + is_mint_rate_limited, + mint_cooldown_remaining, + run_mint_operation, +) from .wallet import ( MintConnectionError, - _is_mint_rate_limited, - _mint_cooldown_remaining, - _mint_operation, get_wallet, is_mint_connection_error, wallet_operation_guard, @@ -31,7 +34,31 @@ lightning_router = APIRouter(prefix="/lightning") # Avoid duplicate work within one process. Cross-process credit fencing is done # by the conditional pending -> paid update in _finalize_invoice_settlement(). -_invoice_settlement_locks: dict[str, asyncio.Lock] = {} +@dataclass +class _InvoiceLockEntry: + lock: asyncio.Lock + users: int = 0 + + +_invoice_settlement_locks: dict[str, _InvoiceLockEntry] = {} + + +@asynccontextmanager +async def _invoice_settlement_lock(invoice_id: str) -> AsyncGenerator[None, None]: + """Serialize one invoice and remove its lock after the last waiter leaves.""" + + entry = _invoice_settlement_locks.get(invoice_id) + if entry is None: + entry = _InvoiceLockEntry(asyncio.Lock()) + _invoice_settlement_locks[invoice_id] = entry + entry.users += 1 + try: + async with entry.lock: + yield + finally: + entry.users -= 1 + if entry.users == 0 and _invoice_settlement_locks.get(invoice_id) is entry: + del _invoice_settlement_locks[invoice_id] @dataclass(frozen=True) @@ -142,11 +169,11 @@ async def _request_mint_with_fallback( configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints] candidates = list(dict.fromkeys(configured)) for mint_url in candidates: - cooldown = _mint_cooldown_remaining(mint_url) + cooldown = mint_cooldown_remaining(mint_url) if cooldown > 0: tried.append(f"{mint_url}: cooling down") logger.info( - "Skipping rate-limited mint", + "Skipping mint during cooldown", extra={ "mint_url": mint_url, "cooldown_seconds": round(cooldown, 2), @@ -156,7 +183,7 @@ async def _request_mint_with_fallback( continue try: wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False) - quote = await _mint_operation( + quote = await run_mint_operation( lambda: wallet.request_mint(amount_sats), op_name="request_mint_invoice", mint_url=mint_url, @@ -165,7 +192,7 @@ async def _request_mint_with_fallback( return quote.request, quote.quote, mint_url except Exception as e: tried.append(f"{mint_url}: {type(e).__name__}") - if not is_mint_connection_error(e) and not _is_mint_rate_limited(e): + if not is_mint_connection_error(e) and not is_mint_rate_limited(e): raise logger.warning( "request_mint failed, trying fallback mint", @@ -352,8 +379,7 @@ async def recover_invoice( async def check_invoice_payment( invoice: LightningInvoice, session: AsyncSession ) -> None: - lock = _invoice_settlement_locks.setdefault(invoice.id, asyncio.Lock()) - async with lock, wallet_operation_guard(): + async with _invoice_settlement_lock(invoice.id), wallet_operation_guard(): minted = False try: # Snapshot the row and end the caller's read transaction before any @@ -368,7 +394,7 @@ async def check_invoice_payment( mint_url = settlement.mint_url or settings.primary_mint wallet = await get_wallet(mint_url, "sat") - mint_status = await _mint_operation( + mint_status = await run_mint_operation( lambda: wallet.get_mint_quote(settlement.payment_hash), op_name="get_mint_quote", mint_url=mint_url, @@ -483,7 +509,7 @@ async def _mint_invoice_quote( return try: - await _mint_operation( + await run_mint_operation( lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash), op_name=f"invoice_mint_{invoice.purpose}", mint_url=mint_url, diff --git a/routstr/mint.py b/routstr/mint.py new file mode 100644 index 00000000..a676a320 --- /dev/null +++ b/routstr/mint.py @@ -0,0 +1,338 @@ +"""Shared policy for bounded, rate-aware Cashu mint API operations.""" + +from __future__ import annotations + +import asyncio +import socket +import time +from contextlib import asynccontextmanager +from contextvars import ContextVar +from typing import Any, AsyncGenerator, Awaitable, Callable + +import httpx + +from .core.logging import get_logger +from .core.settings import settings + +logger = get_logger(__name__) + +MINT_TRANSPORT_EXCEPTIONS: tuple[type[BaseException], ...] = ( + httpx.NetworkError, + httpx.TimeoutException, + ConnectionError, + socket.gaierror, + asyncio.TimeoutError, +) + +MINT_TRANSPORT_COOLDOWN_SECONDS = 30.0 +_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS = 60.0 +_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS = 7 * 60 * 60 + +_fail_fast_depth: ContextVar[int] = ContextVar("mint_fail_fast_depth", default=0) + + +class MintRateLimitedError(httpx.HTTPStatusError): + """Typed boundary error preserving a Cashu mint's HTTP 429 response.""" + + +class MintCooldownError(Exception): + """A mint is cooling down and this operation must not wait.""" + + def __init__(self, mint_url: str, retry_after_seconds: float): + self.mint_url = mint_url + self.retry_after_seconds = max(0.0, retry_after_seconds) + super().__init__( + f"Mint {mint_url} is cooling down; retry after " + f"{self.retry_after_seconds:.2f}s" + ) + + +@asynccontextmanager +async def fail_fast_mint_operations() -> AsyncGenerator[None, None]: + """Make mint cooldown/probe waits fail fast in the current task. + + Wallet mutation code holds a process-wide file lock. It enters this scope so + an existing mint cooldown can never turn that lock into a multi-hour wait. + """ + + token = _fail_fast_depth.set(_fail_fast_depth.get() + 1) + try: + yield + finally: + _fail_fast_depth.reset(token) + + +class MintRateGuard: + """Limit concurrency and remember per-mint cooldown/probe state.""" + + _guards: dict[str, "MintRateGuard"] = {} + + @classmethod + def get(cls, mint_url: str) -> "MintRateGuard": + concurrency = settings.mint_max_concurrency + guard = cls._guards.get(mint_url) + if guard is None or guard._max_concurrency != concurrency: + guard = cls(mint_url, concurrency) + cls._guards[mint_url] = guard + return guard + + def __init__(self, mint_url: str, max_concurrency: int): + self._mint_url = mint_url + self._max_concurrency = max_concurrency + self._semaphore = ( + asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None + ) + self._cooldown_until = 0.0 + self._cooldown_reason: str | None = None + self._consecutive_rate_limits = 0 + self._needs_probe = False + self._probe_lock = asyncio.Lock() + + def apply_cooldown(self, delay: float, *, reason: str | None = None) -> None: + deadline = time.monotonic() + max(0.0, delay) + if deadline >= self._cooldown_until: + self._cooldown_until = deadline + if reason is not None: + self._cooldown_reason = reason + elif self._cooldown_reason is None and reason is not None: + self._cooldown_reason = reason + self._needs_probe = True + + def apply_rate_limit_cooldown(self, retry_after: float | None = None) -> float: + remaining = self.cooldown_remaining() + if remaining > 0 and self._cooldown_reason == "rate_limited": + minimum = min( + _MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, + max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0), + ) + if minimum > remaining: + self.apply_cooldown(minimum, reason="rate_limited") + return minimum + return remaining + + self._consecutive_rate_limits += 1 + base = max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0) + multiplier = 2 ** min(self._consecutive_rate_limits - 1, 10) + delay = min(_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, base * multiplier) + self.apply_cooldown(delay, reason="rate_limited") + return delay + + def cooldown_remaining(self) -> float: + return max(0.0, self._cooldown_until - time.monotonic()) + + def cooldown_reason(self) -> str | None: + return self._cooldown_reason if self.cooldown_remaining() > 0 else None + + def _raise_if_wait_forbidden(self) -> None: + if _fail_fast_depth.get() and ( + self._needs_probe or self.cooldown_remaining() > 0 + ): + raise MintCooldownError(self._mint_url, self.cooldown_remaining()) + + async def _wait_for_cooldown(self) -> None: + while True: + self._raise_if_wait_forbidden() + deadline = self._cooldown_until + wait = max(0.0, deadline - time.monotonic()) + if wait <= 0: + return + logger.debug( + "Mint rate guard: cooling down", + extra={"mint_url": self._mint_url, "wait_seconds": round(wait, 2)}, + ) + await asyncio.sleep(wait) + if self._cooldown_until <= deadline: + return + + async def _run_probe(self, factory: Callable[[], Awaitable[Any]]) -> Any: + await self._wait_for_cooldown() + logger.info( + "Mint cooldown ended; sending one probe request", + extra={"event": "mint_cooldown_probe_started", "mint_url": self._mint_url}, + ) + try: + result = await factory() + except Exception as error: + if is_mint_rate_limited(error): + retry_after = None + if isinstance(error, httpx.HTTPStatusError): + retry_after = parse_retry_after(error.response.headers) + delay = max( + _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, + retry_after or 0.0, + ) + self.apply_cooldown(delay, reason="rate_limited") + else: + self.apply_cooldown(1.0) + logger.warning( + "Mint cooldown probe failed", + extra={ + "event": "mint_cooldown_probe_failed", + "mint_url": self._mint_url, + "error": str(error), + "error_type": type(error).__name__, + "cooldown_seconds": round(self.cooldown_remaining(), 2), + "consecutive_rate_limits": self._consecutive_rate_limits, + }, + ) + raise + + self._needs_probe = False + self._cooldown_until = 0.0 + self._cooldown_reason = None + self._consecutive_rate_limits = 0 + logger.info( + "Mint cooldown probe succeeded; restoring normal concurrency", + extra={ + "event": "mint_cooldown_probe_succeeded", + "mint_url": self._mint_url, + }, + ) + return result + + async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any: + while True: + self._raise_if_wait_forbidden() + if self._needs_probe or self.cooldown_remaining() > 0: + async with self._probe_lock: + self._raise_if_wait_forbidden() + if self.cooldown_remaining() > 0: + self._needs_probe = True + if self._needs_probe: + return await self._run_probe(factory) + continue + + if self._semaphore is None: + return await factory() + async with self._semaphore: + self._raise_if_wait_forbidden() + if self._needs_probe: + continue + return await factory() + + +def mint_cooldown_remaining(mint_url: str) -> float: + return MintRateGuard.get(mint_url).cooldown_remaining() + + +def mint_cooldown_reason(mint_url: str) -> str | None: + return MintRateGuard.get(mint_url).cooldown_reason() + + +def is_mint_rate_limited(error: BaseException) -> bool: + """Return whether an exception chain represents HTTP 429/cooldown.""" + + current: BaseException | None = error + seen: set[int] = set() + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, MintCooldownError): + return True + if isinstance(current, httpx.HTTPStatusError): + if current.response.status_code == 429: + return True + current = current.__cause__ or current.__context__ + return False + + +def parse_retry_after(headers: Any) -> float | None: + raw = headers.get("retry-after") or headers.get("Retry-After") + if raw is None: + return None + try: + return float(str(raw).strip()) + except (TypeError, ValueError): + return None + + +async def run_mint_operation( + factory: Callable[[], Awaitable[Any]], + *, + op_name: str = "mint_operation", + mint_url: str = "", + retry_timeouts: bool = True, + retry_on_rate_limit: bool = True, +) -> Any: + """Run one mint operation with bounded concurrency and adaptive cooldown.""" + + guard = MintRateGuard.get(mint_url) if mint_url else None + timeout = settings.mint_operation_timeout_seconds + max_attempts = settings.mint_retry_max_attempts + 1 + + async def timed_factory() -> Any: + if timeout > 0: + return await asyncio.wait_for(factory(), timeout=timeout) + return await factory() + + async def invoke() -> Any: + if guard is not None: + return await guard.run(timed_factory) + return await timed_factory() + + for attempt in range(max_attempts): + try: + return await invoke() + except MintCooldownError: + raise + except (asyncio.TimeoutError, httpx.TimeoutException) as exc: + if retry_timeouts and attempt < max_attempts - 1: + backoff = (2**attempt) + (time.monotonic() % 1.0) + logger.warning( + "Mint operation timed out, retrying", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "backoff_seconds": round(backoff, 2), + }, + ) + await asyncio.sleep(backoff) + continue + raise httpx.TimeoutException( + f"{op_name} timed out (attempts: {attempt + 1})" + ) from exc + except Exception as exc: + if not is_mint_rate_limited(exc): + raise + + backoff = (2**attempt) + (time.monotonic() % 1.0) + if isinstance(exc, httpx.HTTPStatusError): + retry_after = parse_retry_after(exc.response.headers) + if retry_after is not None: + backoff = max(retry_after, backoff) + cooldown = backoff + if guard is not None: + cooldown = guard.apply_rate_limit_cooldown(backoff) + + if not retry_on_rate_limit: + logger.warning( + "Mint rate-limited, skipping retries for fallback", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "cooldown_seconds": round(cooldown, 2), + "consecutive_rate_limits": guard._consecutive_rate_limits + if guard is not None + else attempt + 1, + }, + ) + raise + + if attempt >= max_attempts - 1: + raise + logger.warning( + "Mint rate-limited, applying cooldown", + extra={ + "op_name": op_name, + "mint_url": mint_url, + "attempt": attempt + 1, + "cooldown_seconds": round(cooldown, 2), + "consecutive_rate_limits": guard._consecutive_rate_limits + if guard is not None + else attempt + 1, + }, + ) + if guard is None: + await asyncio.sleep(cooldown) + + raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 67a7284d..a3ab1fb0 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -18,18 +18,6 @@ from ..wallet import deserialize_token_from_string logger = get_logger(__name__) -# Interim policy: when Routstr must move value to another trusted mint, the -# cross-mint Lightning round trip can consume fees that are not visible to the -# client. Reserve 5% headroom until the fee-payer policy is made explicit. -_MINT_FEE_ALLOWANCE = 0.05 - - -def apply_mint_fee_allowance(cost_msat: int) -> int: - """Reserve headroom for possible trusted-mint fallback fees.""" - adjusted = math.ceil(cost_msat * (1 - _MINT_FEE_ALLOWANCE)) - return max(settings.min_request_msat, adjusted) - - def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None: if x_cashu := headers.get("x-cashu", None): cashu_token = x_cashu diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index fbd28586..c6c3ccb7 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -1,17 +1,13 @@ from __future__ import annotations -import asyncio import math from typing import TypedDict import httpx +from cashu.core.base import MeltQuoteState from cashu.wallet.wallet import Proof, Wallet -# The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung or -# very slow mint can block a melt (and any caller, e.g. the payout loop) -# indefinitely. _mint_operation (imported lazily in raw_send_to_lnurl to avoid -# a circular import with wallet.py) bounds it via MINT_OPERATION_TIMEOUT_SECONDS. -MELT_TIMEOUT_SECONDS = 60 +from ..mint import MINT_TRANSPORT_EXCEPTIONS, run_mint_operation try: from bech32 import bech32_decode, convertbits # type: ignore @@ -222,9 +218,7 @@ async def raw_send_to_lnurl( lnurl_data["callback_url"], final_amount ) - from ..wallet import _mint_operation - - melt_quote_resp = await _mint_operation( + melt_quote_resp = await run_mint_operation( lambda: wallet.melt_quote(invoice=bolt11_invoice), op_name="lnurl_melt_quote", mint_url=str(wallet.url), @@ -234,22 +228,44 @@ async def raw_send_to_lnurl( proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) try: - _ = await asyncio.wait_for( - _mint_operation( - lambda: wallet.melt( - proofs=proofs, - invoice=bolt11_invoice, - fee_reserve_sat=melt_quote_resp.fee_reserve, - quote_id=melt_quote_resp.quote, - ), - op_name="lnurl_melt", - mint_url=str(wallet.url), - retry_timeouts=False, + melt_response = await run_mint_operation( + lambda: wallet.melt( + proofs=proofs, + invoice=bolt11_invoice, + fee_reserve_sat=melt_quote_resp.fee_reserve, + quote_id=melt_quote_resp.quote, ), - timeout=MELT_TIMEOUT_SECONDS, + op_name="lnurl_melt", + mint_url=str(wallet.url), + retry_timeouts=False, ) - except (httpx.TimeoutException, asyncio.TimeoutError) as e: + except MINT_TRANSPORT_EXCEPTIONS as error: + melt_response = None + melt_error: BaseException | None = error + else: + melt_error = None + + if getattr(melt_response, "state", None) == MeltQuoteState.paid: + return final_amount + + try: + quote = await run_mint_operation( + lambda: wallet.get_melt_quote(melt_quote_resp.quote), + op_name="reconcile_lnurl_melt_quote", + mint_url=str(wallet.url), + retry_timeouts=False, + ) + except Exception as reconciliation_error: raise LNURLError( - f"Melt timed out after {MELT_TIMEOUT_SECONDS}s (mint unresponsive)" - ) from e - return final_amount + "Melt outcome is ambiguous; quote reconciliation failed and proofs " + "must not be retried" + ) from reconciliation_error + + if quote is not None and quote.state == MeltQuoteState.paid: + return final_amount + + state = getattr(getattr(quote, "state", None), "value", "unknown") + raise LNURLError( + "Melt outcome is ambiguous; proofs must not be retried " + f"(quote_state={state})" + ) from melt_error diff --git a/routstr/proxy.py b/routstr/proxy.py index 2bbe10ee..9b7d4077 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -27,7 +27,6 @@ from .core.db import ( from .core.exceptions import UpstreamError from .core.not_found import build_not_found_response from .payment.helpers import ( - apply_mint_fee_allowance, calculate_discounted_max_cost, check_token_balance, create_error_response, @@ -354,7 +353,6 @@ async def _proxy( max_cost_for_model = await calculate_discounted_max_cost( _max_cost_for_model, request_body_dict, model_obj=model_obj ) - max_cost_for_model = apply_mint_fee_allowance(max_cost_for_model) check_token_balance(headers, request_body_dict, max_cost_for_model) @@ -493,9 +491,6 @@ async def _proxy( candidate_max = await calculate_discounted_max_cost( candidate_max, request_body_dict, model_obj=model_obj ) - # Apply the same interim 5% trusted-mint fee headroom used for the - # first candidate; failover must not silently change admission. - candidate_max = apply_mint_fee_allowance(candidate_max) if candidate_max > max_cost_for_model: await revert_pay_for_request( key, session, max_cost_for_model, reservation_snapshot diff --git a/routstr/wallet.py b/routstr/wallet.py index a58d4c95..c533355e 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -2,16 +2,15 @@ import asyncio import fcntl import os import re -import socket import time import typing from contextlib import asynccontextmanager from contextvars import ContextVar from pathlib import Path -from typing import Any, AsyncGenerator, Awaitable, Callable, TypedDict +from typing import AsyncGenerator, TypedDict import httpx -from cashu.core.base import MintQuote, Proof, Token +from cashu.core.base import MeltQuoteState, MintQuote, Proof, Token from cashu.core.mint_info import MintInfo as _CashuMintInfo from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet as _CashuWallet @@ -21,8 +20,27 @@ from sqlmodel import col, select, update from .core import db, get_logger from .core.db import store_cashu_transaction_with_retry as store_cashu_transaction from .core.settings import settings +from .mint import ( + MINT_TRANSPORT_COOLDOWN_SECONDS, + MINT_TRANSPORT_EXCEPTIONS, + MintRateGuard, + MintRateLimitedError, + fail_fast_mint_operations, + is_mint_rate_limited, + mint_cooldown_reason, + mint_cooldown_remaining, + run_mint_operation, +) from .payment.lnurl import raw_send_to_lnurl +# Backwards-compatible aliases for callers/tests that imported the former +# wallet-local policy. Production modules use the public routstr.mint API. +_MintRateGuard = MintRateGuard +_mint_operation = run_mint_operation +_mint_cooldown_remaining = mint_cooldown_remaining +_mint_cooldown_reason = mint_cooldown_reason +_is_mint_rate_limited = is_mint_rate_limited + # cashu still declares Optional[X] without explicit defaults on MintInfo. # Under pydantic v2 those are required, but real mints omit many of them. # Default Optional fields to None at import time so balance fetches don't 422. @@ -69,7 +87,8 @@ async def wallet_operation_guard() -> AsyncGenerator[None, None]: except BlockingIOError: await _scheduler_sleep(0.05) depth_token = _wallet_operation_depth.set(1) - yield + async with fail_fast_mint_operations(): + yield finally: if depth_token is not None: _wallet_operation_depth.reset(depth_token) @@ -94,6 +113,20 @@ def _mints_to_inspect() -> list[str]: return mint_urls +class Wallet(_CashuWallet): + """Cashu adapter that preserves HTTP 429 for Routstr's mint policy.""" + + @staticmethod + def raise_on_error_request(resp: httpx.Response) -> None: + if resp.status_code == 429: + raise MintRateLimitedError( + "Cashu mint rate limited", + request=resp.request, + response=resp, + ) + _CashuWallet.raise_on_error_request(resp) + + class MintConnectionError(Exception): """The mint could not be reached (network transport failure). @@ -116,339 +149,6 @@ class TokenConsumedError(Exception): """ -class MintRateLimitedError(httpx.HTTPStatusError): - """Typed boundary error preserving a Cashu mint's HTTP 429 response.""" - - -class Wallet(_CashuWallet): - """Cashu wallet adapter that preserves rate-limit status information. - - Cashu's default response adapter converts JSON error bodies into plain - ``Exception`` instances before calling ``raise_for_status``. Intercept 429 - here so Routstr's fallback and cooldown policy can use the real status - without unreliable message matching. - """ - - @staticmethod - def raise_on_error_request(resp: httpx.Response) -> None: - if resp.status_code == 429: - raise MintRateLimitedError( - "Cashu mint rate limited", - request=resp.request, - response=resp, - ) - _CashuWallet.raise_on_error_request(resp) - - -# httpx base classes cover their subclasses. HTTPStatusError is excluded on -# purpose — that means the mint answered, just with an error status. -_MINT_TRANSPORT_COOLDOWN_SECONDS = 30.0 -_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS = 60.0 -_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS = 7 * 60 * 60 - -_TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( - httpx.NetworkError, - httpx.TimeoutException, - ConnectionError, # refused/reset/aborted - socket.gaierror, # DNS failure - asyncio.TimeoutError, -) - - -class _MintRateGuard: - """Limit concurrency and remember per-mint rate-limit cooldowns.""" - - _guards: dict[str, "_MintRateGuard"] = {} - - @classmethod - def get(cls, mint_url: str) -> "_MintRateGuard": - concurrency = settings.mint_max_concurrency - guard = cls._guards.get(mint_url) - if guard is None or guard._max_concurrency != concurrency: - guard = cls(mint_url, concurrency) - cls._guards[mint_url] = guard - return guard - - def __init__(self, mint_url: str, max_concurrency: int): - self._mint_url = mint_url - self._max_concurrency = max_concurrency - self._semaphore = ( - asyncio.Semaphore(max_concurrency) if max_concurrency > 0 else None - ) - self._cooldown_until = 0.0 - self._cooldown_reason: str | None = None - self._consecutive_rate_limits = 0 - self._needs_probe = False - self._probe_lock = asyncio.Lock() - - def apply_cooldown(self, delay: float, *, reason: str | None = None) -> None: - deadline = time.monotonic() + max(0.0, delay) - if deadline >= self._cooldown_until: - self._cooldown_until = deadline - if reason is not None: - self._cooldown_reason = reason - elif self._cooldown_reason is None and reason is not None: - self._cooldown_reason = reason - self._needs_probe = True - - def apply_rate_limit_cooldown(self, retry_after: float | None = None) -> float: - remaining = self.cooldown_remaining() - if remaining > 0 and self._cooldown_reason == "rate_limited": - minimum = min( - _MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, - max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0), - ) - if minimum > remaining: - self.apply_cooldown(minimum, reason="rate_limited") - return minimum - return remaining - - self._consecutive_rate_limits += 1 - base = max(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, retry_after or 0.0) - multiplier = 2 ** min(self._consecutive_rate_limits - 1, 10) - delay = min(_MINT_RATE_LIMIT_MAX_COOLDOWN_SECONDS, base * multiplier) - self.apply_cooldown(delay, reason="rate_limited") - return delay - - def cooldown_remaining(self) -> float: - return max(0.0, self._cooldown_until - time.monotonic()) - - def cooldown_reason(self) -> str | None: - return self._cooldown_reason if self.cooldown_remaining() > 0 else None - - async def _wait_for_cooldown(self) -> None: - while True: - deadline = self._cooldown_until - wait = max(0.0, deadline - time.monotonic()) - if wait <= 0: - return - logger.debug( - "Mint rate guard: cooling down", - extra={"mint_url": self._mint_url, "wait_seconds": round(wait, 2)}, - ) - await asyncio.sleep(wait) - if self._cooldown_until <= deadline: - return - - async def _run_probe(self, factory: Callable[[], Awaitable[Any]]) -> Any: - await self._wait_for_cooldown() - logger.warning( - "Mint cooldown ended; sending one probe request", - extra={"event": "mint_cooldown_probe_started", "mint_url": self._mint_url}, - ) - try: - result = await factory() - except Exception as error: - # Keep queued callers behind the probe. On a rate-limit, - # re-apply the *same* cooldown the caller already set rather - # than calling apply_rate_limit_cooldown() — the probe is a - # recovery check, not a new request that should escalate the - # exponential backoff counter. - if _is_mint_rate_limited(error): - retry_after = None - if isinstance(error, httpx.HTTPStatusError): - retry_after = _parse_retry_after(error.response.headers) - delay = max( - _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, - retry_after or 0.0, - ) - self.apply_cooldown(delay, reason="rate_limited") - else: - self.apply_cooldown(1.0) - logger.warning( - "Mint cooldown probe failed", - extra={ - "event": "mint_cooldown_probe_failed", - "mint_url": self._mint_url, - "error": str(error), - "error_type": type(error).__name__, - "cooldown_seconds": round(self.cooldown_remaining(), 2), - "consecutive_rate_limits": self._consecutive_rate_limits, - }, - ) - raise - - self._needs_probe = False - self._cooldown_until = 0.0 - self._cooldown_reason = None - self._consecutive_rate_limits = 0 - logger.warning( - "Mint cooldown probe succeeded; restoring normal concurrency", - extra={ - "event": "mint_cooldown_probe_succeeded", - "mint_url": self._mint_url, - }, - ) - return result - - async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any: - while True: - if self._needs_probe or self.cooldown_remaining() > 0: - async with self._probe_lock: - if self.cooldown_remaining() > 0: - self._needs_probe = True - if self._needs_probe: - return await self._run_probe(factory) - continue - - if self._semaphore is None: - return await factory() - async with self._semaphore: - if self._needs_probe: - continue - return await factory() - - -def _mint_cooldown_remaining(mint_url: str) -> float: - return _MintRateGuard.get(mint_url).cooldown_remaining() - - -def _mint_cooldown_reason(mint_url: str) -> str | None: - return _MintRateGuard.get(mint_url).cooldown_reason() - - -def _is_mint_rate_limited(error: BaseException) -> bool: - """True if the mint returned an HTTP 429 (Too Many Requests). - - Only matches ``httpx.HTTPStatusError`` with status code 429 — never - classifies based on the exception's message text. Substring matching - on ``"rate limit"`` / ``"too many requests"`` was removed because it - catches unrelated errors (e.g. a 503 whose body happens to mention - "database rate exceeded"), which triggers unnecessary exponential - backoff and can block state recovery indefinitely. - """ - current: BaseException | None = error - seen: set[int] = set() - while current is not None and id(current) not in seen: - seen.add(id(current)) - if isinstance(current, httpx.HTTPStatusError): - if current.response.status_code == 429: - return True - current = current.__cause__ or current.__context__ - return False - - -async def _mint_operation( - factory: Callable[[], Awaitable[Any]], - *, - op_name: str = "mint_operation", - mint_url: str = "", - retry_timeouts: bool = True, - retry_on_rate_limit: bool = True, -) -> Any: - """Run a mint operation with bounded concurrency and adaptive cooldown. - - The timeout applies to each network attempt. Queueing, cooldown, and retry - backoff are deliberately outside it so the shipped 60-second 429 cooldown - is not cancelled by the 30-second operation timeout. ``factory`` must - return a fresh coroutine for every retry. - - When ``retry_on_rate_limit`` is False a 429 is not retried in-place — the - cooldown is still applied to the per-mint guard (so subsequent operations on - that mint wait), but the exception is re-raised so the caller (typically - ``_request_mint_with_fallback``) can immediately try a different mint. - """ - guard = _MintRateGuard.get(mint_url) if mint_url else None - timeout = settings.mint_operation_timeout_seconds - max_attempts = settings.mint_retry_max_attempts + 1 - - async def timed_factory() -> Any: - if timeout > 0: - return await asyncio.wait_for(factory(), timeout=timeout) - return await factory() - - async def invoke() -> Any: - if guard is not None: - return await guard.run(timed_factory) - return await timed_factory() - - async def run_with_retries() -> Any: - for attempt in range(max_attempts): - try: - return await invoke() - except (asyncio.TimeoutError, httpx.TimeoutException) as exc: - if retry_timeouts and attempt < max_attempts - 1: - backoff = (2**attempt) + (time.monotonic() % 1.0) - logger.warning( - "Mint operation timed out, retrying", - extra={ - "op_name": op_name, - "mint_url": mint_url, - "attempt": attempt + 1, - "backoff_seconds": round(backoff, 2), - }, - ) - await asyncio.sleep(backoff) - continue - raise httpx.TimeoutException( - f"{op_name} timed out (attempts: {attempt + 1})" - ) from exc - except Exception as exc: - if not _is_mint_rate_limited(exc): - raise - - # Apply cooldown to the guard regardless — even when we're - # about to re-raise for fallback, the guard must remember that - # this mint is rate-limited for future operations. - backoff = (2**attempt) + (time.monotonic() % 1.0) - if isinstance(exc, httpx.HTTPStatusError): - retry_after = _parse_retry_after(exc.response.headers) - if retry_after is not None: - backoff = max(retry_after, backoff) - cooldown = backoff - if guard is not None: - cooldown = guard.apply_rate_limit_cooldown(backoff) - - # When the caller has a fallback strategy (trusted-mint - # list), re-raise immediately so the caller can try the next - # mint instead of waiting through this mint's cooldown. - if not retry_on_rate_limit: - logger.warning( - "Mint rate-limited, skipping retries for fallback", - extra={ - "op_name": op_name, - "mint_url": mint_url, - "cooldown_seconds": round(cooldown, 2), - "consecutive_rate_limits": guard._consecutive_rate_limits - if guard is not None - else attempt + 1, - }, - ) - raise - - if attempt >= max_attempts - 1: - raise - logger.warning( - "Mint rate-limited, applying cooldown", - extra={ - "op_name": op_name, - "mint_url": mint_url, - "attempt": attempt + 1, - "cooldown_seconds": round(cooldown, 2), - "consecutive_rate_limits": guard._consecutive_rate_limits - if guard is not None - else attempt + 1, - }, - ) - if guard is None: - await asyncio.sleep(cooldown) - - raise RuntimeError(f"{op_name}: exhausted retries unexpectedly") - - return await run_with_retries() - - -def _parse_retry_after(headers: Any) -> float | None: - """Parse a Retry-After header (delta-seconds form) into seconds.""" - raw = headers.get("retry-after") or headers.get("Retry-After") - if raw is None: - return None - try: - return float(str(raw).strip()) - except (TypeError, ValueError): - return None - - def is_source_mint_connection_error(error: BaseException) -> bool: seen: set[int] = set() current: BaseException | None = error @@ -475,7 +175,7 @@ def is_mint_connection_error(error: BaseException) -> bool: return False if isinstance(current, MintConnectionError): return True - if isinstance(current, _TRANSPORT_EXC_TYPES): + if isinstance(current, MINT_TRANSPORT_EXCEPTIONS): return True current = current.__cause__ or current.__context__ return False @@ -519,7 +219,7 @@ def classify_redemption_error( "The mint that issued this Cashu token is unreachable; the token cannot be redeemed at another mint", "cashu_source_mint_unreachable", ) - if _is_mint_rate_limited(error): + if is_mint_rate_limited(error): return ( "mint_rate_limited", 503, @@ -605,14 +305,14 @@ async def _redeem_same_mint( drifts insolvent. """ try: - await _mint_operation( + await run_mint_operation( lambda: wallet.load_mint(keyset_id=token_obj.keysets[0]), op_name="redeem_load_mint", mint_url=token_obj.mint, ) wallet.verify_proofs_dleq(token_obj.proofs) input_fees = wallet.get_fees_for_proofs(token_obj.proofs) - await _mint_operation( + await run_mint_operation( lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), op_name="redeem_split", mint_url=token_obj.mint, @@ -655,7 +355,7 @@ async def recieve_token( destinations = list( dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) ) - logger.warning( + logger.info( "Cashu cross-mint swap required", extra={ "event": "cashu_swap_started", @@ -667,7 +367,7 @@ async def recieve_token( ) return await swap_to_trusted_mint(token_obj, wallet) - logger.warning( + logger.info( "Trying same-mint Cashu redemption", extra={ "event": "cashu_same_mint_redemption", @@ -769,12 +469,12 @@ async def find_trusted_mint_with_funds( balances: dict[str, int] = {} for mint_url in candidates: - if _mint_cooldown_remaining(mint_url) > 0: + if mint_cooldown_remaining(mint_url) > 0: continue try: wallet = await get_wallet(mint_url, unit, retry_on_rate_limit=False) except Exception as error: - if is_mint_connection_error(error) or _is_mint_rate_limited(error): + if is_mint_connection_error(error) or is_mint_rate_limited(error): balances[mint_url] = 0 continue raise @@ -819,6 +519,17 @@ def _net_minted_amount(amount_msat: int, token_unit: str, fees: int) -> int: return int(remaining_msat) +def _melt_definitively_failed(error: Exception) -> bool: + """Return whether the mint authoritatively rejected the Lightning payment. + + Cashu releases the reserved proofs for these responses, so the token remains + reusable. Transport failures and unknown errors are deliberately excluded: + after dispatch their payment outcome may still be pending or paid. + """ + message = str(error).strip() + return message.lower() == "could not pay invoice." or "(Code: 20004)" in message + + def _melt_insufficient_shortfall(error: Exception) -> int | None: """ Classify a melt failure: return the observed shortfall (in the token unit) @@ -871,7 +582,7 @@ async def _request_mint_with_fallback( f"Token value is too small after fee deduction or unit conversion." ) candidates = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) - logger.warning( + logger.info( "Trying trusted destination mints", extra={ "event": "cashu_destination_candidates", @@ -883,7 +594,7 @@ async def _request_mint_with_fallback( ) tried: list[str] = [] for candidate_index, mint_url in enumerate(candidates, start=1): - cooldown = _mint_cooldown_remaining(mint_url) + cooldown = mint_cooldown_remaining(mint_url) if cooldown > 0: tried.append(f"{mint_url}: cooling down") logger.warning( @@ -898,7 +609,7 @@ async def _request_mint_with_fallback( }, ) continue - logger.warning( + logger.info( "Trying destination mint", extra={ "event": "cashu_destination_attempt", @@ -917,13 +628,13 @@ async def _request_mint_with_fallback( settings.primary_mint_unit, retry_on_rate_limit=False, ) - quote = await _mint_operation( + quote = await run_mint_operation( lambda: wallet.request_mint(amount), op_name=op_name, mint_url=mint_url, retry_on_rate_limit=False, ) - logger.warning( + logger.info( "Destination mint selected", extra={ "event": "cashu_destination_selected", @@ -937,12 +648,12 @@ async def _request_mint_with_fallback( except Exception as error: tried.append(f"{mint_url}: {type(error).__name__}") connection_failure = is_mint_connection_error(error) - rate_limited = _is_mint_rate_limited(error) + rate_limited = is_mint_rate_limited(error) if not connection_failure and not rate_limited: raise if connection_failure: - _MintRateGuard.get(mint_url).apply_cooldown( - _MINT_TRANSPORT_COOLDOWN_SECONDS, reason="unreachable" + MintRateGuard.get(mint_url).apply_cooldown( + MINT_TRANSPORT_COOLDOWN_SECONDS, reason="unreachable" ) logger.warning( "Destination mint failed", @@ -993,7 +704,7 @@ async def _calculate_swap_amount( if token_mint_url == settings.primary_mint: logger.info( - "swap_to_primary_mint: skipping fee estimation (same mint)", + "swap_to_trusted_mint: skipping fee estimation (same mint)", extra={"minted_amount": receive_amount}, ) return int(receive_amount) @@ -1005,7 +716,7 @@ async def _calculate_swap_amount( # logs. Guard early with full diagnostic context instead. if receive_amount <= 0: logger.error( - "swap_to_primary_mint: receive_amount is zero or negative, cannot estimate fees", + "swap_to_trusted_mint: receive_amount is zero or negative, cannot estimate fees", extra={ "amount_msat": amount_msat, "token_unit": token_unit, @@ -1022,7 +733,7 @@ async def _calculate_swap_amount( ) logger.info( - "swap_to_primary_mint: estimating fees", + "swap_to_trusted_mint: estimating fees", extra={ "dummy_amount": receive_amount, "unit": settings.primary_mint_unit, @@ -1040,7 +751,7 @@ async def _calculate_swap_amount( primary_wallet=primary_wallet, ) stage = "source_fee_quote" - dummy_melt_quote = await _mint_operation( + dummy_melt_quote = await run_mint_operation( lambda: token_wallet.melt_quote(dummy_mint_quote.request), op_name="swap_fee_est_melt_quote", mint_url=token_mint_url, @@ -1055,7 +766,7 @@ async def _calculate_swap_amount( raise ValueError(f"Fees ({total_fees} {token_unit}) exceed token amount") logger.info( - "swap_to_primary_mint: fee estimation result", + "swap_to_trusted_mint: fee estimation result", extra={ "token_amount_sat": _msats_to_sats(amount_msat), "estimated_fee": total_fees, @@ -1105,10 +816,62 @@ async def _calculate_swap_amount( raise ValueError(f"Failed to estimate fees: {e}") from e +async def _reconcile_ambiguous_melt( + wallet: Wallet, quote_id: str, proofs: list[Proof] +) -> bool: + """Confirm a dispatched melt is paid or conservatively mark it ambiguous. + + A PAID quote is authoritative and does not require a proof-state lookup. + Every other immediate snapshot remains unsafe to retry: an in-flight + Lightning payment can still move UNPAID/UNSPENT to PENDING or PAID after the + cancelled HTTP request returns. + """ + try: + quote = await run_mint_operation( + lambda: wallet.get_melt_quote(quote_id), + op_name="reconcile_swap_melt_quote", + mint_url=str(wallet.url), + retry_timeouts=False, + ) + except Exception as error: + raise TokenConsumedError( + "Source melt outcome is unknown; reconciliation required" + ) from error + + if quote is not None and quote.state == MeltQuoteState.paid: + return True + + try: + proof_response = await run_mint_operation( + lambda: wallet.check_proof_state(proofs), + op_name="reconcile_swap_proofs", + mint_url=str(wallet.url), + retry_timeouts=False, + ) + proof_states = [state.state.value for state in proof_response.states] + except Exception: + proof_states = [] + + quote_state = getattr(getattr(quote, "state", None), "value", "unknown") + raise TokenConsumedError( + "Source melt outcome is ambiguous; reconciliation required " + f"(quote_state={quote_state}, proof_states={proof_states})" + ) + + +async def _confirm_melt_paid( + wallet: Wallet, quote_id: str, proofs: list[Proof], response: object +) -> bool: + """Accept a melt response only when PAID is explicit or reconciled.""" + if getattr(response, "state", None) == MeltQuoteState.paid: + return True + return await _reconcile_ambiguous_melt(wallet, quote_id, proofs) + + async def swap_to_trusted_mint( token_obj: Token, token_wallet: Wallet ) -> tuple[int, str, str]: - logger.warning( + logger.info( "Starting Cashu cross-mint swap", extra={ "event": "cashu_swap_started", @@ -1135,7 +898,7 @@ async def swap_to_trusted_mint( # NUT-02 input fee still applies; _redeem_same_mint accounts for it. if token_obj.mint == settings.primary_mint: logger.info( - "swap_to_primary_mint: token already on primary mint, skipping swap", + "swap_to_trusted_mint: token already on primary mint, skipping swap", extra={ "mint": token_obj.mint, "amount": token_amount, @@ -1166,7 +929,7 @@ async def swap_to_trusted_mint( attempt += 1 if minted_amount <= 0: logger.error( - "swap_to_primary_mint: minted_amount is zero or negative before requesting quote", + "swap_to_trusted_mint: minted_amount is zero or negative before requesting quote", extra={ "minted_amount": minted_amount, "attempt": attempt, @@ -1188,7 +951,7 @@ async def swap_to_trusted_mint( primary_wallet=primary_wallet, ) logger.info( - "swap_to_primary_mint: mint quote received", + "swap_to_trusted_mint: mint quote received", extra={ "mint_quote_id": mint_quote.quote, "attempt": attempt, @@ -1196,7 +959,7 @@ async def swap_to_trusted_mint( }, ) - logger.warning( + logger.info( "Requesting melt quote from source mint", extra={ "event": "cashu_source_melt_quote_attempt", @@ -1206,7 +969,7 @@ async def swap_to_trusted_mint( }, ) try: - melt_quote = await _mint_operation( + melt_quote = await run_mint_operation( lambda: token_wallet.melt_quote(mint_quote.request), op_name="swap_melt_quote", mint_url=token_obj.mint, @@ -1232,7 +995,7 @@ async def swap_to_trusted_mint( input_fees = token_wallet.get_fees_for_proofs(token_obj.proofs) total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees logger.info( - "swap_to_primary_mint: melt quote received", + "swap_to_trusted_mint: melt quote received", extra={ "melt_quote_id": melt_quote.quote, "melt_amount": melt_quote.amount, @@ -1252,7 +1015,7 @@ async def swap_to_trusted_mint( ) if attempt >= _MAX_SWAP_ATTEMPTS or recomputed <= 0: logger.warning( - "swap_to_primary_mint: insufficient token amount for melt fees", + "swap_to_trusted_mint: insufficient token amount for melt fees", extra={ "token_amount": token_amount, "melt_amount": melt_quote.amount, @@ -1269,7 +1032,7 @@ async def swap_to_trusted_mint( f"(amount: {melt_quote.amount} + fee: {melt_quote.fee_reserve} + input_fees: {input_fees})" ) logger.warning( - "swap_to_primary_mint: melt quote exceeds token amount, retrying", + "swap_to_trusted_mint: melt quote exceeds token amount, retrying", extra={ "total_needed": total_needed, "token_amount": token_amount, @@ -1281,7 +1044,7 @@ async def swap_to_trusted_mint( continue try: - _ = await _mint_operation( + melt_response = await run_mint_operation( lambda: token_wallet.melt( proofs=token_obj.proofs, invoice=mint_quote.request, @@ -1292,36 +1055,45 @@ async def swap_to_trusted_mint( mint_url=token_obj.mint, retry_timeouts=False, ) + await _confirm_melt_paid( + token_wallet, melt_quote.quote, token_obj.proofs, melt_response + ) except Exception as e: - # A down mint won't fix itself by retrying with a smaller amount. - if is_mint_connection_error(e): - logger.error( - "Source mint became unreachable during melt", - extra={ - "event": "cashu_source_mint_unreachable", - "stage": "source_melt", - "error": str(e), - "error_type": type(e).__name__, - "source_mint": token_obj.mint, - "destination_mint": dest_mint_url, - "attempt": attempt, - }, - ) - raise SourceMintConnectionError( - "Issuing Cashu mint is unreachable" - ) from e shortfall = _melt_insufficient_shortfall(e) - recomputed = 0 - if shortfall is not None: - observed_extra_fee += shortfall - recomputed = _net_minted_amount( - amount_msat, - token_obj.unit, - melt_quote.fee_reserve + input_fees + observed_extra_fee, - ) - if shortfall is None or attempt >= _MAX_SWAP_ATTEMPTS or recomputed <= 0: + if shortfall is None: + if isinstance(e, TokenConsumedError): + raise + if _melt_definitively_failed(e): + raise ValueError( + f"Failed to melt token from foreign mint {token_obj.mint}: {e}" + ) from e + if is_mint_connection_error(e): + await _reconcile_ambiguous_melt( + token_wallet, melt_quote.quote, token_obj.proofs + ) + logger.info( + "Source melt reconciled as paid; minting on destination", + extra={ + "event": "cashu_source_melt_reconciled_paid", + "source_mint": token_obj.mint, + "destination_mint": dest_mint_url, + "melt_quote_id": melt_quote.quote, + }, + ) + break + raise TokenConsumedError( + "Source melt failed after dispatch; outcome requires reconciliation" + ) from e + + observed_extra_fee += shortfall + recomputed = _net_minted_amount( + amount_msat, + token_obj.unit, + melt_quote.fee_reserve + input_fees + observed_extra_fee, + ) + if attempt >= _MAX_SWAP_ATTEMPTS or recomputed <= 0: logger.error( - "swap_to_primary_mint: melt failed", + "swap_to_trusted_mint: melt failed", extra={ "error": str(e), "error_type": type(e).__name__, @@ -1336,7 +1108,7 @@ async def swap_to_trusted_mint( f"Failed to melt token from foreign mint {token_obj.mint}: {e}" ) from e logger.warning( - "swap_to_primary_mint: mint demanded more than quoted at melt, retrying", + "swap_to_trusted_mint: mint demanded more than quoted at melt, retrying", extra={ "shortfall": shortfall, "retry_minted_amount": recomputed, @@ -1348,7 +1120,7 @@ async def swap_to_trusted_mint( break - logger.warning( + logger.info( "Source melt succeeded; minting on destination", extra={ "event": "cashu_destination_mint_attempt", @@ -1361,9 +1133,9 @@ async def swap_to_trusted_mint( await dest_wallet.load_proofs(reload=True) pre_mint_balance = dest_wallet.available_balance.amount try: - _ = await _mint_operation( + _ = await run_mint_operation( lambda: dest_wallet.mint(minted_amount, quote_id=mint_quote.quote), - op_name="swap_mint_on_primary", + op_name="swap_mint_on_destination", mint_url=dest_mint_url, retry_timeouts=False, ) @@ -1373,7 +1145,7 @@ async def swap_to_trusted_mint( # bump_secret_derivation ran locally. Recover orphaned proofs and # advance the counter so the next request derives fresh secrets. logger.warning( - "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", + "swap_to_trusted_mint: outputs already signed — recovering orphaned proofs", extra={ "mint_quote_id": mint_quote.quote, "minted_amount": minted_amount, @@ -1388,7 +1160,7 @@ async def swap_to_trusted_mint( post_recovery_balance = dest_wallet.available_balance.amount balance_gained = post_recovery_balance - pre_mint_balance logger.info( - "swap_to_primary_mint: recovery scan completed", + "swap_to_trusted_mint: recovery scan completed", extra={ "pre_mint_balance": pre_mint_balance, "post_recovery_balance": post_recovery_balance, @@ -1411,7 +1183,7 @@ async def swap_to_trusted_mint( raise except Exception as recovery_err: logger.error( - "swap_to_primary_mint: recovery failed", + "swap_to_trusted_mint: recovery failed", extra={"error": str(recovery_err)}, ) raise TokenConsumedError( @@ -1419,7 +1191,7 @@ async def swap_to_trusted_mint( ) from e else: logger.error( - "swap_to_primary_mint: mint on primary failed after successful melt", + "swap_to_trusted_mint: mint on primary failed after successful melt", extra={ "error": str(e), "error_type": type(e).__name__, @@ -1432,7 +1204,7 @@ async def swap_to_trusted_mint( "Mint on primary failed after successful melt" ) from e - logger.warning( + logger.info( "Cashu cross-mint swap completed", extra={ "event": "cashu_swap_completed", @@ -1595,13 +1367,13 @@ async def get_wallet( now = time.monotonic() last = _wallet_last_load.get(id) if last is None or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: - await _mint_operation( + await run_mint_operation( lambda: _wallets[id].load_mint(), op_name="load_mint", mint_url=mint_url, retry_on_rate_limit=retry_on_rate_limit, ) - await _mint_operation( + await run_mint_operation( lambda: _wallets[id].load_proofs(reload=True), op_name="load_proofs", mint_url=mint_url, @@ -1640,7 +1412,7 @@ async def slow_filter_spend_proofs( batch_size = 1000 for i in range(0, len(proofs), batch_size): pb = proofs[i : i + batch_size] - proof_states = await _mint_operation( + proof_states = await run_mint_operation( lambda: wallet.check_proof_state(pb), op_name="check_proof_state", mint_url=str(wallet.url), @@ -1652,12 +1424,7 @@ async def slow_filter_spend_proofs( else: _spent_proofs.append(proof) if _spent_proofs: - await _mint_operation( - lambda: wallet.set_reserved_for_send(_spent_proofs, reserved=True), - op_name="set_reserved_spent_proofs", - mint_url=str(wallet.url), - retry_timeouts=False, - ) + await wallet.set_reserved_for_send(_spent_proofs, reserved=True) return _proofs @@ -1686,7 +1453,7 @@ async def _get_supported_mint_units(mint_url: str) -> list[str]: return cached[1] wallet = await get_wallet(mint_url, settings.primary_mint_unit, load=False) - keysets = await _mint_operation( + keysets = await run_mint_operation( lambda: wallet._get_keysets(), op_name="get_mint_keysets", mint_url=mint_url, @@ -1750,7 +1517,7 @@ async def fetch_all_balances( mint_units[mint_url] = await _get_supported_mint_units(mint_url) except Exception as error: connection_failure = is_mint_connection_error(error) - rate_limited = _is_mint_rate_limited(error) + rate_limited = is_mint_rate_limited(error) error_code = ( "rate_limited" if rate_limited @@ -1759,12 +1526,12 @@ async def fetch_all_balances( else "mint_error" ) if connection_failure: - _MintRateGuard.get(mint_url).apply_cooldown( + MintRateGuard.get(mint_url).apply_cooldown( _BALANCE_FETCH_RETRY_SECONDS, reason="unreachable" ) retry_delay = max( _BALANCE_FETCH_RETRY_SECONDS, - _mint_cooldown_remaining(mint_url), + mint_cooldown_remaining(mint_url), ) discovery_errors.append( _balance_error( @@ -1819,9 +1586,9 @@ async def fetch_all_balances( retry_after_seconds=failure[0] - now, ) - cooldown = _mint_cooldown_remaining(mint_url) + cooldown = mint_cooldown_remaining(mint_url) if cooldown > 0: - error_code = _mint_cooldown_reason(mint_url) or "cooldown" + error_code = mint_cooldown_reason(mint_url) or "cooldown" error = { "rate_limited": "Mint is rate limited", "unreachable": "Mint is unreachable", @@ -1846,7 +1613,7 @@ async def fetch_all_balances( proofs = await slow_filter_spend_proofs(proofs, wallet) except Exception as error: connection_failure = is_mint_connection_error(error) - rate_limited = _is_mint_rate_limited(error) + rate_limited = is_mint_rate_limited(error) error_code = ( "rate_limited" if rate_limited @@ -1855,16 +1622,16 @@ async def fetch_all_balances( else "mint_error" ) if rate_limited: - _MintRateGuard.get(mint_url).apply_rate_limit_cooldown( + MintRateGuard.get(mint_url).apply_rate_limit_cooldown( _BALANCE_FETCH_RETRY_SECONDS ) elif connection_failure: - _MintRateGuard.get(mint_url).apply_cooldown( + MintRateGuard.get(mint_url).apply_cooldown( _BALANCE_FETCH_RETRY_SECONDS, reason=error_code ) retry_delay = max( _BALANCE_FETCH_RETRY_SECONDS, - _mint_cooldown_remaining(mint_url), + mint_cooldown_remaining(mint_url), ) _balance_fetch_failures[key] = ( time.monotonic() + retry_delay, diff --git a/tests/integration/test_insufficient_balance.py b/tests/integration/test_insufficient_balance.py index 3acc5db9..a63c4251 100644 --- a/tests/integration/test_insufficient_balance.py +++ b/tests/integration/test_insufficient_balance.py @@ -208,25 +208,30 @@ async def test_pay_for_request_succeeds_when_balance_equals_cost( @pytest.mark.asyncio -async def test_five_percent_mint_fallback_headroom_is_admitted_and_reserved( +async def test_full_model_maximum_is_required_and_reserved( integration_session: AsyncSession, ) -> None: from routstr.auth import pay_for_request, validate_bearer_key - from routstr.payment.helpers import apply_mint_fee_allowance - key = _key(balance=95_000) - integration_session.add(key) + short_key = _key(balance=95_000) + exact_key = _key(balance=100_000) + integration_session.add(short_key) + integration_session.add(exact_key) await integration_session.commit() - admission_cost = apply_mint_fee_allowance(100_000) - validated = await validate_bearer_key( - f"sk-{key.hashed_key}", integration_session, min_cost=admission_cost - ) - await pay_for_request(validated, admission_cost, integration_session) + with pytest.raises(HTTPException) as insufficient: + await validate_bearer_key( + f"sk-{short_key.hashed_key}", integration_session, min_cost=100_000 + ) + assert insufficient.value.status_code == 402 - await integration_session.refresh(key) - assert admission_cost == 95_000 - assert key.reserved_balance == 95_000 + validated = await validate_bearer_key( + f"sk-{exact_key.hashed_key}", integration_session, min_cost=100_000 + ) + await pay_for_request(validated, 100_000, integration_session) + + await integration_session.refresh(exact_key) + assert exact_key.reserved_balance == 100_000 # --------------------------------------------------------------------------- @@ -288,7 +293,7 @@ async def test_http_402_response_shape_on_insufficient_balance( error = body["detail"]["error"] assert error["code"] == "insufficient_balance" assert error["type"] == "insufficient_quota" - assert "591.744 sats (591744 msats) required" in error["message"] + assert "622.888 sats (622888 msats) required" in error["message"] assert "20.32 sats (20320 msats) available" in error["message"] # Balance must be completely untouched diff --git a/tests/integration/test_swap_fee_retry.py b/tests/integration/test_swap_fee_retry.py index 138a4d4d..b27360a2 100644 --- a/tests/integration/test_swap_fee_retry.py +++ b/tests/integration/test_swap_fee_retry.py @@ -20,6 +20,7 @@ from collections.abc import Callable from unittest.mock import AsyncMock, Mock, patch import pytest +from cashu.core.base import MeltQuoteState from httpx import AsyncClient, Response from routstr.core.settings import settings @@ -81,7 +82,9 @@ def _make_swap_mocks( quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee() ) ) - mock_token_wallet.melt = AsyncMock(return_value=Mock()) + mock_token_wallet.melt = AsyncMock( + return_value=Mock(state=MeltQuoteState.paid) + ) return mock_token, mock_token_wallet, mock_primary_wallet @@ -144,7 +147,7 @@ async def test_topup_retries_when_melt_demands_more_than_quoted( "Mint Error: not enough inputs provided for melt. " "Provided: 179, needed: 180 (Code: 11000)" ), - Mock(), + Mock(state=MeltQuoteState.paid), ] response = await _post_topup( diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index 7618ed47..ceeca1e4 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -161,7 +161,7 @@ async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: patch.object(settings, "primary_mint", "http://mint:3338"), patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.db.create_session", _fake_session), - patch("routstr.wallet.time.monotonic", return_value=10), + patch("routstr.mint.time.monotonic", return_value=10), patch("routstr.wallet.logger.warning") as warning, ): first = await fetch_all_balances(units=["sat"]) @@ -180,7 +180,7 @@ async def test_fetch_all_balances_backs_off_after_connection_failure() -> None: patch.object(settings, "primary_mint", "http://mint:3338"), patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.db.create_session", _fake_session), - patch("routstr.wallet.time.monotonic", return_value=71), + patch("routstr.mint.time.monotonic", return_value=71), patch("routstr.wallet.logger.warning"), ): await fetch_all_balances(units=["sat"]) @@ -219,7 +219,7 @@ async def test_balance_failure_applies_mint_cooldown_to_other_units() -> None: patch.object(settings, "primary_mint", mint), patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.db.create_session", _fake_session), - patch("routstr.wallet.time.monotonic", return_value=10), + patch("routstr.mint.time.monotonic", return_value=10), patch("routstr.wallet.logger.warning") as warning, ): details, *_ = await fetch_all_balances(units=["sat", "msat"]) diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py index c44954f2..d8741b7e 100644 --- a/tests/unit/test_lightning_settlement.py +++ b/tests/unit/test_lightning_settlement.py @@ -155,6 +155,7 @@ async def test_non_pending_invoice_is_not_minted() -> None: get_wallet.assert_not_awaited() session.commit.assert_awaited_once() + assert _invoice_settlement_locks == {} @pytest.mark.asyncio @@ -212,3 +213,4 @@ async def test_concurrent_invoice_checks_finalize_once_in_process() -> None: assert invoice.status == "paid" finalize.assert_awaited_once() + assert _invoice_settlement_locks == {} diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index 47efbcf2..7311b4a2 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -1,70 +1,127 @@ -"""raw_send_to_lnurl() must not hang forever on an unresponsive mint. - -The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung -mint would block the melt (and the payout loop) indefinitely. raw_send_to_lnurl -now wraps wallet.melt() in asyncio.wait_for(MELT_TIMEOUT_SECONDS) and surfaces a -timeout as LNURLError instead of hanging. -""" +"""LNURL melt attempts must not misclassify ambiguous payment outcomes.""" import asyncio +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest +from cashu.core.base import MeltQuoteState -from routstr.payment import lnurl +from routstr.core.settings import settings from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl +LNURL_DATA = { + "callback_url": "https://ln.tld/cb", + "min_sendable": 1_000, + "max_sendable": 100_000_000, +} -@pytest.mark.asyncio -async def test_raw_send_to_lnurl_times_out_on_hung_melt() -> None: + +def _wallet() -> tuple[MagicMock, list[MagicMock]]: proofs = [MagicMock(amount=1000)] - - wallet = MagicMock() + wallet = MagicMock(url="https://mint.test") wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q")) wallet.select_to_send = AsyncMock(return_value=(proofs, None)) + return wallet, proofs + + +def _lnurl_patches() -> tuple[Any, Any]: + return ( + patch( + "routstr.payment.lnurl.get_lnurl_data", + AsyncMock(return_value=LNURL_DATA), + ), + patch( + "routstr.payment.lnurl.get_lnurl_invoice", + AsyncMock(return_value=("lnbc1...", {})), + ), + ) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> None: + wallet, proofs = _wallet() async def _hang(**kwargs: object) -> None: - await asyncio.sleep(5) # far longer than the patched timeout + await asyncio.sleep(5) wallet.melt = AsyncMock(side_effect=_hang) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + data_patch, invoice_patch = _lnurl_patches() - lnurl_data = { - "callback_url": "https://ln.tld/cb", - "min_sendable": 1_000, - "max_sendable": 100_000_000, - } - - with patch.object(lnurl, "MELT_TIMEOUT_SECONDS", 0.05), patch( - "routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=lnurl_data) - ), patch( - "routstr.payment.lnurl.get_lnurl_invoice", - AsyncMock(return_value=("lnbc1...", {})), + with ( + patch.object(settings, "mint_operation_timeout_seconds", 0.05), + patch.object(settings, "mint_retry_max_attempts", 0), + data_patch, + invoice_patch, + pytest.raises(LNURLError, match="outcome is ambiguous"), ): - with pytest.raises(LNURLError, match="Melt timed out"): - await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.get_melt_quote.assert_awaited_once_with("q") + wallet.set_reserved_for_melt.assert_not_called() @pytest.mark.asyncio -async def test_raw_send_to_lnurl_succeeds_within_timeout() -> None: - """A prompt melt still returns the net amount, unaffected by the guard.""" - proofs = [MagicMock(amount=1000)] +async def test_raw_send_to_lnurl_timeout_reconciled_paid_is_success() -> None: + wallet, proofs = _wallet() - wallet = MagicMock() - wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q")) - wallet.select_to_send = AsyncMock(return_value=(proofs, None)) - wallet.melt = AsyncMock(return_value=MagicMock()) + async def _hang(**kwargs: object) -> None: + await asyncio.sleep(5) - lnurl_data = { - "callback_url": "https://ln.tld/cb", - "min_sendable": 1_000, - "max_sendable": 100_000_000, - } + wallet.melt = AsyncMock(side_effect=_hang) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.paid) + ) + data_patch, invoice_patch = _lnurl_patches() - with patch.object(lnurl, "MELT_TIMEOUT_SECONDS", 5), patch( - "routstr.payment.lnurl.get_lnurl_data", AsyncMock(return_value=lnurl_data) - ), patch( - "routstr.payment.lnurl.get_lnurl_invoice", - AsyncMock(return_value=("lnbc1...", {})), + with ( + patch.object(settings, "mint_operation_timeout_seconds", 0.05), + patch.object(settings, "mint_retry_max_attempts", 0), + data_patch, + invoice_patch, + ): + paid = await raw_send_to_lnurl( + wallet, proofs, "owner@ln.tld", "sat", amount=1000 + ) + + assert paid > 0 + wallet.get_melt_quote.assert_awaited_once_with("q") + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending)) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.pending) + ) + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_operation_timeout_seconds", 5), + data_patch, + invoice_patch, + pytest.raises(LNURLError, match="outcome is ambiguous"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.get_melt_quote.assert_awaited_once_with("q") + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_succeeds_on_explicit_paid_response() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid)) + wallet.get_melt_quote = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_operation_timeout_seconds", 5), + data_patch, + invoice_patch, ): paid = await raw_send_to_lnurl( wallet, proofs, "owner@ln.tld", "sat", amount=1000 @@ -72,3 +129,4 @@ async def test_raw_send_to_lnurl_succeeds_within_timeout() -> None: assert paid > 0 wallet.melt.assert_awaited_once() + wallet.get_melt_quote.assert_not_awaited() diff --git a/tests/unit/test_melt_reconciliation.py b/tests/unit/test_melt_reconciliation.py new file mode 100644 index 00000000..a68cb64b --- /dev/null +++ b/tests/unit/test_melt_reconciliation.py @@ -0,0 +1,91 @@ +from unittest.mock import AsyncMock, Mock + +import pytest +from cashu.core.base import MeltQuoteState, ProofSpentState + +from routstr.wallet import ( + TokenConsumedError, + _confirm_melt_paid, + _reconcile_ambiguous_melt, +) + + +@pytest.mark.asyncio +async def test_paid_quote_is_authoritative_when_proof_lookup_would_fail() -> None: + wallet = Mock( + url="http://source-mint:3338", + get_melt_quote=AsyncMock(return_value=Mock(state=MeltQuoteState.paid)), + check_proof_state=AsyncMock(side_effect=RuntimeError("proof API unavailable")), + ) + + assert await _reconcile_ambiguous_melt(wallet, "quote-1", [Mock()]) is True + wallet.check_proof_state.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_timeout_snapshot_unpaid_unspent_remains_non_retryable() -> None: + wallet = Mock( + url="http://source-mint:3338", + get_melt_quote=AsyncMock(return_value=Mock(state=MeltQuoteState.unpaid)), + check_proof_state=AsyncMock( + return_value=Mock(states=[Mock(state=ProofSpentState.unspent)]) + ), + ) + + with pytest.raises(TokenConsumedError, match="ambiguous"): + await _reconcile_ambiguous_melt(wallet, "quote-2", [Mock()]) + + +@pytest.mark.asyncio +async def test_successful_pending_melt_response_requires_reconciliation() -> None: + wallet = Mock( + url="http://source-mint:3338", + get_melt_quote=AsyncMock(return_value=Mock(state=MeltQuoteState.pending)), + check_proof_state=AsyncMock( + return_value=Mock(states=[Mock(state=ProofSpentState.pending)]) + ), + ) + + with pytest.raises(TokenConsumedError, match="ambiguous"): + await _confirm_melt_paid( + wallet, + "quote-pending", + [Mock()], + Mock(state=MeltQuoteState.pending), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("quote_state", "proof_state"), + [ + (MeltQuoteState.pending, ProofSpentState.pending), + (MeltQuoteState.unpaid, ProofSpentState.spent), + (MeltQuoteState.unpaid, ProofSpentState.pending), + ], +) +async def test_ambiguous_or_consumed_melt_is_never_reported_unspent( + quote_state: MeltQuoteState, proof_state: ProofSpentState +) -> None: + wallet = Mock( + url="http://source-mint:3338", + get_melt_quote=AsyncMock(return_value=Mock(state=quote_state)), + check_proof_state=AsyncMock( + return_value=Mock(states=[Mock(state=proof_state)]) + ), + ) + + with pytest.raises(TokenConsumedError, match="reconciliation required"): + await _reconcile_ambiguous_melt(wallet, "quote-3", [Mock()]) + + +@pytest.mark.asyncio +async def test_failed_melt_reconciliation_is_non_retryable() -> None: + wallet = Mock( + url="http://source-mint:3338", + get_melt_quote=AsyncMock(side_effect=RuntimeError("mint unavailable")), + check_proof_state=AsyncMock(), + ) + + with pytest.raises(TokenConsumedError, match="outcome is unknown"): + await _reconcile_ambiguous_melt(wallet, "quote-4", [Mock()]) diff --git a/tests/unit/test_mint.py b/tests/unit/test_mint.py new file mode 100644 index 00000000..5a258a29 --- /dev/null +++ b/tests/unit/test_mint.py @@ -0,0 +1,65 @@ +from unittest.mock import AsyncMock, Mock, patch + +import httpx +import pytest +from cashu.core.base import Unit + +from routstr.mint import ( + MintCooldownError, + MintRateGuard, + MintRateLimitedError, + fail_fast_mint_operations, +) +from routstr.wallet import Wallet + + +@pytest.mark.asyncio +async def test_cooldown_fails_fast_while_wallet_mutation_scope_is_held() -> None: + guard = MintRateGuard("http://mint:3338", max_concurrency=1) + guard.apply_cooldown(3600, reason="rate_limited") + operation = AsyncMock(return_value="should not run") + + with ( + patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, + pytest.raises(MintCooldownError) as caught, + ): + async with fail_fast_mint_operations(): + await guard.run(operation) + + assert caught.value.retry_after_seconds > 0 + operation.assert_not_awaited() + sleep.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cashu_429_dispatches_through_wallet_override() -> None: + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 429, + request=request, + json={"detail": "too many requests", "code": 42900}, + ) + + wallet = object.__new__(Wallet) + wallet.url = "http://mint:3338" + wallet.db = Mock() + wallet.keysets = {"loaded": Mock()} + wallet.mint_info = Mock() + wallet.mint_info.requires_blind_auth_path.return_value = False + wallet.mint_info.requires_clear_auth_path.return_value = False + wallet.auth_db = None + wallet.auth_keyset_id = None + + real_client = httpx.AsyncClient + + def client_factory(*args: object, **kwargs: object) -> httpx.AsyncClient: + return real_client( + transport=httpx.MockTransport(handler), + base_url=str(kwargs["base_url"]), + ) + + with ( + patch("cashu.wallet.v1_api.httpx.AsyncClient", side_effect=client_factory), + pytest.raises(MintRateLimitedError), + ): + await wallet.mint_quote(1, Unit.sat) diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index fe0cd573..ef8dde63 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -7,21 +7,7 @@ os.environ["UPSTREAM_BASE_URL"] = "http://test" os.environ["UPSTREAM_API_KEY"] = "test" from routstr.core.settings import settings # noqa: E402 -from routstr.payment.helpers import ( # noqa: E402 - apply_mint_fee_allowance, - get_max_cost_for_model, -) - - -def test_mint_fee_allowance_reserves_five_percent_fallback_headroom() -> None: - # Interim policy: Routstr may pay hidden cross-mint Lightning fees when a - # trusted-mint fallback is required. - assert apply_mint_fee_allowance(124_886) == 118_642 - - -def test_mint_fee_allowance_never_drops_below_minimum() -> None: - with patch.object(settings, "min_request_msat", 100): - assert apply_mint_fee_allowance(50) == 100 +from routstr.payment.helpers import get_max_cost_for_model # noqa: E402 async def test_get_max_cost_for_model_known() -> None: diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 512dfa31..31dd5767 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -422,4 +422,4 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: with pytest.raises(asyncio.CancelledError): await proxy_module.proxy(request, "v1/chat/completions", session=session) - revert_mock.assert_awaited_once_with(key, session, 950, reservation_snapshot) + revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot) diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index 65b50a73..495f1e57 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -406,4 +406,4 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: assert RAW_ORG_ID not in serialized assert "org-[REDACTED]" in serialized # Single upstream failed -> reservation reverted exactly once (no double-charge). - revert_mock.assert_awaited_once_with(key, session, 950, reservation) + revert_mock.assert_awaited_once_with(key, session, 1000, reservation) diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 257b6132..03123a60 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, Mock, patch import httpx import pytest +from cashu.core.base import MeltQuoteState from routstr.core.db import ApiKey from routstr.wallet import ( @@ -639,7 +640,9 @@ def _make_swap_mocks( quote=f"melt_quote_{invoice}", amount=invoice, fee_reserve=_next_fee() ) ) - mock_token_wallet.melt = AsyncMock(return_value=Mock()) + mock_token_wallet.melt = AsyncMock( + return_value=Mock(state=MeltQuoteState.paid) + ) return mock_token, mock_token_wallet, mock_primary_wallet @@ -770,7 +773,7 @@ async def test_swap_retries_when_melt_demands_more_than_quoted() -> None: "Mint Error: not enough inputs provided for melt. " "Provided: 179, needed: 180 (Code: 11000)" ), - Mock(), + Mock(state=MeltQuoteState.paid), ] from routstr.core.settings import settings @@ -801,7 +804,7 @@ async def test_swap_retries_on_cdk_unbalanced_error() -> None: ) mock_token_wallet.melt.side_effect = [ Exception("Mint Error: Transaction unbalanced: 179, 178, 2 (Code: 11005)"), - Mock(), + Mock(state=MeltQuoteState.paid), ] from routstr.core.settings import settings @@ -1522,22 +1525,32 @@ async def test_swap_fee_estimation_transport_error_raises_mint_connection_error( @pytest.mark.asyncio -async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: - """A transport failure during melt is surfaced as MintConnectionError and - is NOT retried — the mint is down, not demanding higher fees.""" +async def test_swap_melt_transport_error_is_never_reported_reusable() -> None: + """A timed-out melt remains ambiguous even when an immediate snapshot says + UNPAID/UNSPENT, so callers must not receive the original token for retry.""" from routstr.wallet import swap_to_primary_mint mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks( 1000, fee_reserves=[10, 10] ) mock_token_wallet.melt = AsyncMock(side_effect=httpx.ConnectTimeout("timed out")) + from cashu.core.base import MeltQuoteState, ProofSpentState + + mock_token_wallet.get_melt_quote = AsyncMock( + return_value=Mock(state=MeltQuoteState.unpaid) + ) + mock_token_wallet.check_proof_state = AsyncMock( + return_value=Mock( + states=[Mock(state=ProofSpentState.unspent) for _ in mock_token.proofs] + ) + ) from routstr.core.settings import settings with patch.object(settings, "primary_mint", "http://primary:3338"): with patch.object(settings, "primary_mint_unit", "sat"): with patch("routstr.wallet.get_wallet", return_value=mock_primary_wallet): - with pytest.raises(MintConnectionError): + with pytest.raises(TokenConsumedError, match="ambiguous"): await swap_to_primary_mint(mock_token, mock_token_wallet) assert mock_token_wallet.melt.call_count == 1 @@ -1595,8 +1608,8 @@ async def test_mint_rate_guard_waits_for_adaptive_cooldown() -> None: guard._cooldown_until = 15.0 operation = AsyncMock(return_value="ok") - with patch("routstr.wallet.time.monotonic", return_value=10.0): - with patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep: + with patch("routstr.mint.time.monotonic", return_value=10.0): + with patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep: assert await guard.run(operation) == "ok" sleep.assert_awaited_once_with(5.0) @@ -1611,7 +1624,7 @@ async def test_mint_rate_guard_exponentially_backs_off_repeated_429s() -> None: expected_delays = [60, 120, 240, 480, 960, 1920, 3840, 7680, 15360, 25200] now = 0.0 - with patch("routstr.wallet.time.monotonic") as monotonic: + with patch("routstr.mint.time.monotonic") as monotonic: for index, expected in enumerate(expected_delays, start=1): monotonic.return_value = now assert guard.apply_rate_limit_cooldown(60) == expected @@ -1682,8 +1695,8 @@ async def test_mint_rate_guard_keeps_cooldown_when_concurrency_is_unlimited() -> operation = AsyncMock(return_value="ok") with ( patch.object(settings, "mint_max_concurrency", 0), - patch("routstr.wallet.time.monotonic", return_value=0), - patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + patch("routstr.mint.time.monotonic", return_value=0), + patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, ): guard = _MintRateGuard.get("http://mint:3338") guard.apply_cooldown(5) @@ -1715,8 +1728,8 @@ async def test_mint_operation_honors_retry_after_as_minimum() -> None: with patch.object(settings, "mint_retry_max_attempts", 1): with patch.object(settings, "mint_operation_timeout_seconds", 0): with patch.object(settings, "mint_max_concurrency", 1): - with patch("routstr.wallet.time.monotonic", return_value=0.1): - with patch("routstr.wallet.asyncio.sleep", sleep): + with patch("routstr.mint.time.monotonic", return_value=0.1): + with patch("routstr.mint.asyncio.sleep", sleep): result = await _mint_operation( factory, mint_url="http://mint:3338" ) @@ -1734,7 +1747,7 @@ async def test_mint_operation_timeout_excludes_adaptive_cooldown() -> None: with ( patch.object(settings, "mint_max_concurrency", 1), patch.object(settings, "mint_operation_timeout_seconds", 0.01), - patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, ): guard = _MintRateGuard.get("http://mint:3338") guard.apply_cooldown(60) @@ -1763,7 +1776,7 @@ async def test_default_timeout_allows_retry_after_rate_limit_cooldown() -> None: patch.object(settings, "mint_retry_max_attempts", 3), patch.object(settings, "mint_operation_timeout_seconds", 30), patch.object(settings, "mint_max_concurrency", 1), - patch("routstr.wallet.asyncio.sleep", AsyncMock()), + patch("routstr.mint.asyncio.sleep", AsyncMock()), ): assert await _mint_operation(operation, mint_url="http://mint:3338") == "ok" @@ -1780,7 +1793,7 @@ async def test_mint_operation_retries_httpx_timeout_only_when_safe() -> None: with patch.object(settings, "mint_retry_max_attempts", 2): with patch.object(settings, "mint_operation_timeout_seconds", 0): - with patch("routstr.wallet.asyncio.sleep", AsyncMock()): + with patch("routstr.mint.asyncio.sleep", AsyncMock()): assert await _mint_operation(retrying) == "ok" with pytest.raises(httpx.TimeoutException): await _mint_operation(non_retrying, retry_timeouts=False) @@ -1802,7 +1815,7 @@ async def test_get_wallet_initializes_and_loads_once_concurrently() -> None: ) as create: # A fresh wallet must load even when the host has been up for less than # the reload interval. - with patch("routstr.wallet.time.monotonic", return_value=10.0): + with patch("routstr.mint.time.monotonic", return_value=10.0): first, second = await asyncio.gather( get_wallet("http://mint:3338"), get_wallet("http://mint:3338") ) @@ -1833,7 +1846,7 @@ async def test_get_wallet_can_surface_429_without_retrying() -> None: patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=wallet)), patch.object(settings, "mint_retry_max_attempts", 3), patch.object(settings, "mint_operation_timeout_seconds", 0), - patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, ): with pytest.raises(httpx.HTTPStatusError): await get_wallet("http://mint:3338", retry_on_rate_limit=False) @@ -1965,7 +1978,7 @@ async def test_swap_falls_back_when_primary_wallet_cannot_load() -> None: melt_quote=AsyncMock( return_value=Mock(quote="melt_q", amount=990, fee_reserve=10) ), - melt=AsyncMock(return_value=Mock()), + melt=AsyncMock(return_value=Mock(state=MeltQuoteState.paid)), ) mint_quote = Mock(quote="mint_q_secondary", request="lnbc1secondary") @@ -1994,6 +2007,7 @@ async def test_swap_falls_back_when_primary_wallet_cannot_load() -> None: patch("asyncio.sleep", AsyncMock()), patch("routstr.wallet.get_wallet", side_effect=mock_get), patch("routstr.wallet.logger.warning") as warning, + patch("routstr.wallet.logger.info") as info, ): amount, unit, mint_url = await swap_to_primary_mint(token, source_wallet) @@ -2003,7 +2017,7 @@ async def test_swap_falls_back_when_primary_wallet_cannot_load() -> None: assert any(call.args[0] == secondary for call in mock_get.await_args_list) events = { call.kwargs["extra"]["event"] - for call in warning.call_args_list + for call in [*warning.call_args_list, *info.call_args_list] if "extra" in call.kwargs and "event" in call.kwargs["extra"] } assert "cashu_destination_failed" in events @@ -2181,8 +2195,8 @@ async def test_wallet_fallback_skips_mint_during_cooldown() -> None: patch.object(settings, "cashu_mints", [primary, secondary]), patch.object(settings, "mint_max_concurrency", 0), patch.object(settings, "mint_operation_timeout_seconds", 0), - patch("routstr.wallet.time.monotonic", return_value=10), - patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep, + patch("routstr.mint.time.monotonic", return_value=10), + patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, patch( "routstr.wallet.get_wallet", AsyncMock(side_effect=lambda mint, *args, **kwargs: wallets[mint]), @@ -2378,49 +2392,28 @@ def test_classify_500_with_rate_limit_text_is_not_mint_rate_limited() -> None: @pytest.mark.asyncio async def test_probe_does_not_escalate_consecutive_rate_limits() -> None: - """When a probe fails with a rate limit, _consecutive_rate_limits should - NOT increment — the probe is a recovery check, not a new request.""" - from routstr.wallet import _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, _MintRateGuard + from routstr.mint import MintRateGuard - guard = _MintRateGuard("http://mint", max_concurrency=0) - - # Simulate initial rate limit: apply_rate_limit_cooldown increments counter + guard = MintRateGuard("http://mint", max_concurrency=0) guard.apply_rate_limit_cooldown() - assert guard._consecutive_rate_limits == 1 - cooldown_before = guard._cooldown_until - assert cooldown_before > 0 + guard._cooldown_until = 0.0 - # Simulate probe failure: _run_probe uses apply_cooldown, NOT - # apply_rate_limit_cooldown, so the counter stays at 1. - guard.apply_cooldown(_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, reason="rate_limited") - assert guard._consecutive_rate_limits == 1 # unchanged! + with pytest.raises(httpx.HTTPStatusError): + await guard.run(AsyncMock(side_effect=_http_429_error())) + + assert guard._consecutive_rate_limits == 1 assert guard._needs_probe is True @pytest.mark.asyncio async def test_probe_recovery_resets_consecutive_rate_limits() -> None: - """A successful probe resets _consecutive_rate_limits to 0.""" - from routstr.wallet import _MintRateGuard + from routstr.mint import MintRateGuard - guard = _MintRateGuard("http://mint", max_concurrency=0) - - # First rate limit: increments to 1, sets 60s cooldown. + guard = MintRateGuard("http://mint", max_concurrency=0) guard.apply_rate_limit_cooldown() - assert guard._consecutive_rate_limits == 1 - - # Manually expire the cooldown so the next call creates a fresh one. guard._cooldown_until = 0.0 - guard._cooldown_reason = None - # Second rate limit (after cooldown expired): increments to 2. - guard.apply_rate_limit_cooldown() - assert guard._consecutive_rate_limits == 2 - - # Simulate a successful probe by resetting (as _run_probe does) - guard._needs_probe = False - guard._cooldown_until = 0.0 - guard._cooldown_reason = None - guard._consecutive_rate_limits = 0 + assert await guard.run(AsyncMock(return_value="ok")) == "ok" assert guard._consecutive_rate_limits == 0 assert guard._needs_probe is False From 19236ecc9db82a05dcaec6a4f988888bf5e5b5b5 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Fri, 31 Jul 2026 02:14:50 +0200 Subject: [PATCH 42/46] fix(db): increase default connection pool capacity --- .env.example | 6 +++--- routstr/core/settings.py | 14 +++++++------- tests/unit/test_settings.py | 12 ++++++------ 3 files changed, 16 insertions(+), 16 deletions(-) diff --git a/.env.example b/.env.example index 9f0c0bbc..17ae6951 100644 --- a/.env.example +++ b/.env.example @@ -26,9 +26,9 @@ ROUTSTR_SECRET_KEY= # logged at startup. Keep total capacity across all workers below the database # connection limit. Pre-ping is automatic for networked backends; SQLite may # explicitly opt in if desired. -# DATABASE_POOL_SIZE=5 -# DATABASE_MAX_OVERFLOW=10 -# DATABASE_POOL_TIMEOUT=30 +# DATABASE_POOL_SIZE=10 +# DATABASE_MAX_OVERFLOW=20 +# DATABASE_POOL_TIMEOUT=15 # DATABASE_POOL_RECYCLE=1800 # DATABASE_POOL_PRE_PING=false # Warn when a checkout is held this many seconds. diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 0caffd16..4674733a 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -106,14 +106,14 @@ class Settings(BaseSettings): default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS" ) - # Database connection-pool controls (advanced). Capacity defaults match - # SQLAlchemy's established queue-pool behavior. Pre-ping is enabled by the - # engine factory for networked backends; SQLite can explicitly opt in. - # These fields are env-only below. - database_pool_size: int = Field(default=5, ge=1, env="DATABASE_POOL_SIZE") - database_max_overflow: int = Field(default=10, ge=0, env="DATABASE_MAX_OVERFLOW") + # Database connection-pool controls (advanced). Capacity defaults provide + # headroom for Routstr's concurrent request and background-payment workload. + # Pre-ping is enabled by the engine factory for networked backends; SQLite + # can explicitly opt in. These fields are env-only below. + database_pool_size: int = Field(default=10, ge=1, env="DATABASE_POOL_SIZE") + database_max_overflow: int = Field(default=20, ge=0, env="DATABASE_MAX_OVERFLOW") database_pool_timeout: float = Field( - default=30.0, gt=0, env="DATABASE_POOL_TIMEOUT" + default=15.0, gt=0, env="DATABASE_POOL_TIMEOUT" ) database_pool_recycle: int = Field(default=1800, ge=0, env="DATABASE_POOL_RECYCLE") database_pool_pre_ping: bool = Field(default=False, env="DATABASE_POOL_PRE_PING") diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index 98a9e24f..34fa7101 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -62,11 +62,11 @@ def test_payout_settings_have_sensible_defaults() -> None: assert s.payout_interval_seconds == 900 -def test_database_pool_defaults_match_sqlalchemy_capacity() -> None: +def test_database_pool_defaults_provide_concurrency_headroom() -> None: s = Settings() - assert s.database_pool_size == 5 - assert s.database_max_overflow == 10 - assert s.database_pool_timeout == 30.0 + assert s.database_pool_size == 10 + assert s.database_max_overflow == 20 + assert s.database_pool_timeout == 15.0 assert s.database_pool_recycle == 1800 assert s.database_pool_pre_ping is False assert s.database_pool_hold_warn_seconds == 10.0 @@ -139,7 +139,7 @@ async def test_update_does_not_apply_env_only_fields_to_live_settings( from the running pool. """ monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False) - monkeypatch.setattr(settings, "database_pool_size", 5) + monkeypatch.setattr(settings, "database_pool_size", 10) engine = create_async_engine("sqlite+aiosqlite:///:memory:") async with AsyncSession(engine, expire_on_commit=False) as session: @@ -151,7 +151,7 @@ async def test_update_does_not_apply_env_only_fields_to_live_settings( # A non-env-only field still updates normally... assert settings.name == "PoolTweaker" # ...but the env-only pool size stays at the boot value. - assert settings.database_pool_size == 5 + assert settings.database_pool_size == 10 # ...and it is never written to the settings blob. blob = await _read_settings_blob(session) assert "database_pool_size" not in blob From e903aa3a9f37210b36e6851654d6b72a94d9a926 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 2 Aug 2026 23:16:01 +0200 Subject: [PATCH 43/46] clean up --- routstr/upstream/model_paths.py | 34 ++++++++++++++++++++++----------- tests/unit/test_model_paths.py | 21 +++++++------------- 2 files changed, 30 insertions(+), 25 deletions(-) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 020f0fbb..95b6edf2 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -31,6 +31,8 @@ from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_sessio from ..core.logging import get_logger if TYPE_CHECKING: + from sqlmodel.ext.asyncio.session import AsyncSession + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -782,18 +784,28 @@ async def get_all_model_paths() -> dict: async def get_paths_for_model(model_id: str) -> dict: - """Return paths only for the exact model ID advertised by ``/v1/models``.""" - async with create_session() as session: - rows = ( - await session.exec( - select(ModelPathRow) - .where(col(ModelPathRow.model_id) == model_id) - .order_by( - col(ModelPathRow.path), - col(ModelPathRow.upstream_provider_id), + """Return paths for an advertised ID or its provider-prefixed alias.""" + + async def load_rows(session: AsyncSession, lookup_id: str) -> list[ModelPathRow]: + return list( + ( + await session.exec( + select(ModelPathRow) + .where(col(ModelPathRow.model_id) == lookup_id) + .order_by( + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) ) - ) - ).all() + ).all() + ) + + async with create_session() as session: + rows = await load_rows(session, model_id) + if not rows: + unprefixed_id = public_model_id(model_id) + if unprefixed_id != model_id: + rows = await load_rows(session, unprefixed_id) seen: set[str] = set() paths: list[dict] = [] diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 1bbe25e0..240a4199 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -462,9 +462,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") == { - _expected_path(1, "shared-model") - } + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} @pytest.mark.asyncio @@ -498,12 +496,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") == { - _expected_path(1, "shared-model") - } - assert _paths_of(payload, "private-alias") == { - _expected_path(2, "private-alias") - } + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} + assert _paths_of(payload, "private-alias") == {_expected_path(2, "private-alias")} @pytest.mark.asyncio @@ -1134,7 +1128,7 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( @pytest.mark.asyncio -async def test_get_paths_for_model_requires_exact_advertised_id( +async def test_get_paths_for_model_accepts_provider_prefixed_alias( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -1158,15 +1152,14 @@ async def test_get_paths_for_model_requires_exact_advertised_id( _path_entry(4, "deepseek-v4-pro"), _path_entry(7, "deepseek-v4-pro"), ] - assert prefixed_paths == [] + assert prefixed_paths == short_paths @pytest.mark.asyncio async def test_get_paths_for_model_multi_segment_id_matches_models_listing( patched_session: AsyncEngine, ) -> None: - """For three-segment ids the discovery id must be the same base id the - rest of the system exposes (first-slash rule), not the last segment.""" + """Three-segment upstream IDs resolve to the same first-slash public ID.""" provider = _FakeProvider( provider_type="generic", base_url="https://x/v1", @@ -1181,7 +1174,7 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" - ] == [] + ] == [_path_entry(1, "fireworks/models/glm-5")] # --------------------------------------------------------------------------- # From dd8c4a9a8aa27b712f91a230ead6b64cf103ce4c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 3 Aug 2026 00:05:36 +0200 Subject: [PATCH 44/46] update migartion --- ..._table.py => 64ed5594df1f_add_model_paths_table.py} | 10 +++++----- ...ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py} | 10 +++++----- tests/unit/test_mint_url_migration.py | 4 ++-- 3 files changed, 12 insertions(+), 12 deletions(-) rename migrations/versions/{4f2a4f3f62e0_add_model_paths_table.py => 64ed5594df1f_add_model_paths_table.py} (91%) rename migrations/versions/{bf76270b66c4_add_mint_url_to_lightning_invoices.py => ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py} (72%) diff --git a/migrations/versions/4f2a4f3f62e0_add_model_paths_table.py b/migrations/versions/64ed5594df1f_add_model_paths_table.py similarity index 91% rename from migrations/versions/4f2a4f3f62e0_add_model_paths_table.py rename to migrations/versions/64ed5594df1f_add_model_paths_table.py index 59fee690..a957cc8d 100644 --- a/migrations/versions/4f2a4f3f62e0_add_model_paths_table.py +++ b/migrations/versions/64ed5594df1f_add_model_paths_table.py @@ -1,8 +1,8 @@ """add model paths table -Revision ID: 4f2a4f3f62e0 -Revises: bf76270b66c4 -Create Date: 2026-08-02 23:28:24.760061 +Revision ID: 64ed5594df1f +Revises: aa50fde387a2 +Create Date: 2026-08-02 22:26:33.280409 """ import sqlalchemy as sa @@ -10,8 +10,8 @@ import sqlmodel from alembic import op # revision identifiers, used by Alembic. -revision = "4f2a4f3f62e0" -down_revision = "bf76270b66c4" +revision = "64ed5594df1f" +down_revision = "aa50fde387a2" branch_labels = None depends_on = None diff --git a/migrations/versions/bf76270b66c4_add_mint_url_to_lightning_invoices.py b/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py similarity index 72% rename from migrations/versions/bf76270b66c4_add_mint_url_to_lightning_invoices.py rename to migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py index f78851fe..21c6e547 100644 --- a/migrations/versions/bf76270b66c4_add_mint_url_to_lightning_invoices.py +++ b/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py @@ -1,16 +1,16 @@ """add mint url to lightning invoices -Revision ID: bf76270b66c4 -Revises: aa50fde387a2 -Create Date: 2026-07-30 00:54:30.306876 +Revision ID: ecfa0d6e2a36 +Revises: 64ed5594df1f +Create Date: 2026-08-02 23:53:00.037456 """ import sqlalchemy as sa from alembic import op # revision identifiers, used by Alembic. -revision = "bf76270b66c4" -down_revision = "aa50fde387a2" +revision = "ecfa0d6e2a36" +down_revision = "64ed5594df1f" branch_labels = None depends_on = None diff --git a/tests/unit/test_mint_url_migration.py b/tests/unit/test_mint_url_migration.py index 398e2ca5..b83566f1 100644 --- a/tests/unit/test_mint_url_migration.py +++ b/tests/unit/test_mint_url_migration.py @@ -32,12 +32,12 @@ def test_mint_url_migration_upgrades_and_downgrades_from_main_head( root = Path(__file__).resolve().parents[2] database_path = tmp_path / "mint-url-migration.db" database_url = f"sqlite+aiosqlite:///{database_path}" - previous_head = "aa50fde387a2" + previous_head = "64ed5594df1f" _run_alembic(root, database_url, "upgrade", previous_head) assert "mint_url" not in _lightning_invoice_columns(database_path) - _run_alembic(root, database_url, "upgrade", "bf76270b66c4") + _run_alembic(root, database_url, "upgrade", "ecfa0d6e2a36") assert "mint_url" in _lightning_invoice_columns(database_path) _run_alembic(root, database_url, "downgrade", previous_head) From da859f2f8419fb6f97fc99008c9aecd7d6a199f6 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 3 Aug 2026 01:42:44 +0200 Subject: [PATCH 45/46] clean up --- ...2a36_add_mint_url_to_lightning_invoices.py | 42 +++ routstr/auth.py | 16 +- routstr/core/admin.py | 42 +-- routstr/core/db.py | 11 +- routstr/lightning.py | 205 ++++++++--- routstr/mint.py | 23 +- routstr/payment/helpers.py | 2 +- routstr/payment/lnurl.py | 16 +- routstr/wallet.py | 325 +++++++++++++----- tests/integration/conftest.py | 9 +- .../test_lightning_invoice_constraints.py | 8 +- .../integration/test_lightning_settlement.py | 94 ++++- .../test_periodic_payout_safety.py | 6 +- tests/integration/test_prune_dead_api_keys.py | 9 +- tests/integration/test_swap_fee_retry.py | 4 +- tests/unit/test_admin_withdraw.py | 114 ++++-- tests/unit/test_auth_cashu.py | 48 +++ tests/unit/test_coverage_admin.py | 25 +- tests/unit/test_fee_payout_crash_safety.py | 2 +- tests/unit/test_lightning_settlement.py | 162 ++++++++- tests/unit/test_lnurl_melt_timeout.py | 72 ++++ tests/unit/test_mint.py | 56 +++ tests/unit/test_payment_helpers.py | 30 ++ tests/unit/test_wallet.py | 259 +++++++++++++- ui/lib/api/services/wallet.ts | 1 + 25 files changed, 1342 insertions(+), 239 deletions(-) diff --git a/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py b/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py index 21c6e547..7d8abffc 100644 --- a/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py +++ b/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py @@ -5,6 +5,9 @@ Revises: 64ed5594df1f Create Date: 2026-08-02 23:53:00.037456 """ +import json +import os + import sqlalchemy as sa from alembic import op @@ -15,11 +18,50 @@ branch_labels = None depends_on = None +def _resolve_backfill_mint_url(bind: sa.engine.Connection) -> str | None: + """Best-effort resolution of the mint that issued pre-existing invoices. + + Order: persisted settings JSON -> PRIMARY_MINT_URL env -> first CASHU_MINTS entry. + """ + try: + row = bind.execute( + sa.text("SELECT data FROM settings ORDER BY id LIMIT 1") + ).fetchone() + if row and row[0]: + data = json.loads(row[0]) + mint = data.get("primary_mint") or next( + iter(data.get("cashu_mints") or []), None + ) + if mint: + return str(mint) + except Exception: + pass + + env_mint = os.environ.get("PRIMARY_MINT_URL", "").strip() + if env_mint: + return env_mint + + cashu_mints = os.environ.get("CASHU_MINTS", "").strip() + if cashu_mints: + return cashu_mints.split(",")[0].strip() or None + return None + + def upgrade() -> None: op.add_column( "lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True) ) + bind = op.get_bind() + backfill_mint = _resolve_backfill_mint_url(bind) + if backfill_mint: + bind.execute( + sa.text( + "UPDATE lightning_invoices SET mint_url = :mint WHERE mint_url IS NULL" + ), + {"mint": backfill_mint}, + ) + def downgrade() -> None: op.drop_column("lightning_invoices", "mint_url") diff --git a/routstr/auth.py b/routstr/auth.py index 610469d1..588fc2d1 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -387,11 +387,23 @@ async def _validate_bearer_key_locked( "has_expiry_time": bool(key_expiry_time), }, ) - if token_obj.mint in settings.cashu_mints: + if token_obj.mint == settings.primary_mint: + if token_obj.unit != settings.primary_mint_unit: + raise redemption_error_to_http_exception( + ValueError( + "Cashu token unit does not match the configured primary " + f"mint unit: expected {settings.primary_mint_unit}, " + f"got {token_obj.unit}" + ) + ) + refund_currency = token_obj.unit + refund_mint_url = settings.primary_mint + elif token_obj.mint in settings.cashu_mints: refund_currency = token_obj.unit refund_mint_url = token_obj.mint else: - refund_currency = "sat" + # Foreign tokens are swapped into the configured primary mint. + refund_currency = settings.primary_mint_unit refund_mint_url = settings.primary_mint new_key = ApiKey( diff --git a/routstr/core/admin.py b/routstr/core/admin.py index f7883261..0b03cdc3 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -13,13 +13,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession from ..payment.models import _row_to_model, list_models from ..proxy import refresh_model_maps, reinitialize_upstreams -from ..wallet import ( - fetch_all_balances, - get_proofs_per_mint_and_unit, - get_wallet, - send_token, - slow_filter_spend_proofs, -) +from ..wallet import fetch_all_balances, send_token, token_mint_url from . import vault from .db import ( ApiKey, @@ -442,37 +436,31 @@ class WithdrawRequest(BaseModel): async def withdraw( request: Request, withdraw_request: WithdrawRequest ) -> dict[str, str]: - # Get wallet and check balance from .settings import settings as global_settings effective_mint = withdraw_request.mint_url or global_settings.primary_mint - wallet = await get_wallet(effective_mint, withdraw_request.unit) - proofs = get_proofs_per_mint_and_unit( - wallet, - effective_mint, - withdraw_request.unit, - not_reserved=True, - ) - proofs = await slow_filter_spend_proofs(proofs, wallet) - current_balance = sum(proof.amount for proof in proofs) - if withdraw_request.amount <= 0: raise HTTPException( status_code=400, detail="Withdrawal amount must be positive" ) - if withdraw_request.amount > current_balance: - raise HTTPException(status_code=400, detail="Insufficient wallet balance") - - token = await send_token( - withdraw_request.amount, withdraw_request.unit, effective_mint - ) + try: + token = await send_token( + withdraw_request.amount, withdraw_request.unit, effective_mint + ) + except ValueError as error: + if not str(error).startswith("No trusted mint has "): + raise + raise HTTPException( + status_code=400, detail="Insufficient wallet balance" + ) from error + actual_mint = token_mint_url(token, effective_mint) try: await store_cashu_transaction( token=token, amount=withdraw_request.amount, unit=withdraw_request.unit, - mint_url=effective_mint, + mint_url=actual_mint, typ="out", collected=False, source="admin", @@ -483,10 +471,10 @@ async def withdraw( extra={ "amount": withdraw_request.amount, "unit": withdraw_request.unit, - "mint_url": effective_mint, + "mint_url": actual_mint, }, ) - return {"token": token} + return {"token": token, "mint_url": actual_mint} class ModelCreate(BaseModel): diff --git a/routstr/core/db.py b/routstr/core/db.py index 4185791d..d9207d64 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -293,7 +293,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in """Delete dead parentless API keys; return the count removed. Dead = 0 balance/reservation/spend/requests, older than the grace period, - no parent, no children, no pending invoice. Cashu rows are unlinked (not + no parent, no children, no retryable invoice. Cashu rows are unlinked (not deleted) first to keep the audit trail. """ cutoff = int(time.time()) - min_age_seconds @@ -307,7 +307,9 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in pending_invoice = ( select(LightningInvoice.id) .where(col(LightningInvoice.api_key_hash) == col(ApiKey.hashed_key)) - .where(col(LightningInvoice.status) == "pending") + .where( + col(LightningInvoice.status).in_(("pending", "settlement_pending")) + ) ).exists() eligible_hashes = ( @@ -435,7 +437,10 @@ class LightningInvoice(SQLModel, table=True): # type: ignore payment_hash: str = Field(description="Payment hash for tracking", unique=True) status: str = Field( default="pending", - description="pending, paid, expired, cancelled, reconciliation_required", + description=( + "pending, settlement_pending, paid, expired, cancelled, " + "reconciliation_required" + ), ) api_key_hash: str | None = Field( default=None, description="Associated API key hash for topup operations" diff --git a/routstr/lightning.py b/routstr/lightning.py index 41dc775f..7a610c10 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -7,6 +7,7 @@ from contextlib import asynccontextmanager from dataclasses import dataclass from typing import Any, AsyncGenerator +from cashu.core.base import MintQuoteState from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel, Field from sqlalchemy.orm.attributes import set_committed_value @@ -32,8 +33,8 @@ logger = get_logger(__name__) lightning_router = APIRouter(prefix="/lightning") -# Avoid duplicate work within one process. Cross-process credit fencing is done -# by the conditional pending -> paid update in _finalize_invoice_settlement(). +# Avoid duplicate work within one process. Cross-process settlement is fenced +# by claiming a paid quote before minting and by the final conditional update. @dataclass class _InvoiceLockEntry: lock: asyncio.Lock @@ -138,6 +139,9 @@ class InvoiceStatusResponse(BaseModel): expires_at: int +_RETRYABLE_INVOICE_STATUSES = ("pending", "settlement_pending") + + class InvoiceRecoverRequest(BaseModel): bolt11: str = Field(description="BOLT11 invoice string") @@ -166,8 +170,11 @@ async def _request_mint_with_fallback( f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}." ) tried: list[str] = [] - configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints] - candidates = list(dict.fromkeys(configured)) + candidates = ( + list(dict.fromkeys(allowed_mints)) + if allowed_mints + else _trusted_mint_candidates() + ) for mint_url in candidates: cooldown = mint_cooldown_remaining(mint_url) if cooldown > 0: @@ -317,12 +324,12 @@ async def get_invoice_status( if not invoice: raise HTTPException(status_code=404, detail="Invoice not found") - if invoice.status == "pending": - await check_invoice_payment(invoice, session) - - if invoice.status == "pending" and int(time.time()) > invoice.expires_at: - invoice.status = "expired" - await session.commit() + definitively_unpaid = False + if invoice.status in _RETRYABLE_INVOICE_STATUSES: + definitively_unpaid = await check_invoice_payment(invoice, session) + await _expire_invoice_if_authoritatively_unpaid( + invoice, session, definitively_unpaid + ) api_key = None if invoice.status == "paid" and invoice.purpose == "create": @@ -356,8 +363,12 @@ async def recover_invoice( if not invoice: raise HTTPException(status_code=404, detail="Invoice not found") - if invoice.status == "pending": - await check_invoice_payment(invoice, session) + definitively_unpaid = False + if invoice.status in _RETRYABLE_INVOICE_STATUSES: + definitively_unpaid = await check_invoice_payment(invoice, session) + await _expire_invoice_if_authoritatively_unpaid( + invoice, session, definitively_unpaid + ) api_key = None if invoice.status == "paid": @@ -376,19 +387,58 @@ async def recover_invoice( ) +async def _claim_paid_invoice_for_settlement( + invoice: LightningInvoice, + caller_session: AsyncSession, + observed_status: str, +) -> bool: + """Claim an authoritative paid quote before consuming it at the mint.""" + if observed_status == "settlement_pending": + return True + if observed_status != "pending": + await _reload_invoice_view(invoice, caller_session) + return False + + async with create_session() as claim_session: + claim = await claim_session.exec( # type: ignore[call-overload] + update(LightningInvoice) + .where( + col(LightningInvoice.id) == invoice.id, + col(LightningInvoice.status) == "pending", + ) + .values(status="settlement_pending") + .execution_options(synchronize_session=False) + ) + await claim_session.commit() + + if claim.rowcount != 1: + await _reload_invoice_view(invoice, caller_session) + return False + + _publish_invoice_value(invoice, "status", "settlement_pending") + return True + + async def check_invoice_payment( invoice: LightningInvoice, session: AsyncSession -) -> None: +) -> bool: + """Settle an invoice and report whether its quote is definitively unpaid. + + False covers paid, pending, and ambiguous transport/DB outcomes so callers + never expire a quote merely because reconciliation could not complete. + """ async with _invoice_settlement_lock(invoice.id), wallet_operation_guard(): minted = False + payment_confirmed = False try: # Snapshot the row and end the caller's read transaction before any # potentially slow mint I/O. All final DB mutations use owned, # short-lived sessions below. await session.refresh(invoice) - if invoice.status != "pending": + if invoice.status not in _RETRYABLE_INVOICE_STATUSES: await session.commit() - return + return False + observed_status = invoice.status settlement = _InvoiceSettlement.from_invoice(invoice) await session.commit() @@ -400,7 +450,15 @@ async def check_invoice_payment( mint_url=mint_url, ) if not mint_status.paid: - return + return getattr(mint_status, "state", None) == MintQuoteState.unpaid + payment_confirmed = True + + # Fence expiry and other workers before consuming the paid quote. + # If a concurrent expiry/finalization won, this worker must not mint. + if not await _claim_paid_invoice_for_settlement( + invoice, session, observed_status + ): + return False # Reject a paid top-up whose target was pruned before redeeming its # single-use quote. The validation session is closed before mint I/O. @@ -416,7 +474,9 @@ async def check_invoice_payment( update(LightningInvoice) .where( col(LightningInvoice.id) == settlement.id, - col(LightningInvoice.status) == "pending", + col(LightningInvoice.status).in_( + _RETRYABLE_INVOICE_STATUSES + ), ) .values(status="reconciliation_required") ) @@ -431,7 +491,7 @@ async def check_invoice_payment( "Paid topup invoice target API key was not found; reconciliation required", extra={"invoice_id": settlement.id}, ) - return + return False # Quote-linked proof verification makes an ambiguous mint response # retryable without crediting unrelated wallet balance growth. @@ -445,7 +505,7 @@ async def check_invoice_payment( ) if not settled: await _reload_invoice_view(invoice, session) - return + return False _publish_invoice_value(invoice, "status", "paid") _publish_invoice_value(invoice, "paid_at", paid_at) @@ -461,9 +521,33 @@ async def check_invoice_payment( else None, }, ) + return False except BaseException as error: # Never roll back the caller-owned session: doing so expires invoice # and sibling ORM objects. Owned sessions roll themselves back. + if payment_confirmed and invoice.status != "settlement_pending": + try: + async with create_session() as state_session: + pending = await state_session.exec( # type: ignore[call-overload] + update(LightningInvoice) + .where( + col(LightningInvoice.id) == invoice.id, + col(LightningInvoice.status).in_( + _RETRYABLE_INVOICE_STATUSES + ), + ) + .values(status="settlement_pending") + ) + await state_session.commit() + if pending.rowcount == 1: + _publish_invoice_value( + invoice, "status", "settlement_pending" + ) + except Exception as state_error: + logger.critical( + "Paid invoice reconciliation state could not be persisted", + extra={"invoice_id": invoice.id, "error": str(state_error)}, + ) if minted: logger.critical( "Invoice mint succeeded but DB finalization failed; reconciliation required", @@ -476,6 +560,7 @@ async def check_invoice_payment( if not isinstance(error, Exception): raise logger.error(f"Failed to check invoice payment: {error}") + return False def _is_outputs_already_signed(error: BaseException) -> bool: @@ -588,7 +673,9 @@ async def _finalize_invoice_settlement( claim = await session.exec( # type: ignore[call-overload] update(LightningInvoice) .where(col(LightningInvoice.id) == invoice.id) - .where(col(LightningInvoice.status) == "pending") + .where( + col(LightningInvoice.status).in_(_RETRYABLE_INVOICE_STATUSES) + ) .values(status="paid", paid_at=paid_at, api_key_hash=api_key_hash) .execution_options(synchronize_session=False) ) @@ -623,6 +710,39 @@ async def _reload_invoice_view( _publish_invoice_value(invoice, "api_key_hash", api_key_hash) +async def _expire_invoice_if_authoritatively_unpaid( + invoice: LightningInvoice, + caller_session: AsyncSession, + definitively_unpaid: bool, +) -> bool: + """Expire one overdue unpaid invoice without overwriting concurrent settlement.""" + if ( + not definitively_unpaid + or invoice.status != "pending" + or int(time.time()) <= invoice.expires_at + ): + return False + + async with create_session() as expiry_session: + expired = await expiry_session.exec( # type: ignore[call-overload] + update(LightningInvoice) + .where( + col(LightningInvoice.id) == invoice.id, + col(LightningInvoice.status) == "pending", + ) + .values(status="expired") + .execution_options(synchronize_session=False) + ) + await expiry_session.commit() + + if expired.rowcount == 1: + _publish_invoice_value(invoice, "status", "expired") + return True + + await _reload_invoice_view(invoice, caller_session) + return False + + async def _credit_topup_record( invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession ) -> None: @@ -635,32 +755,33 @@ INVOICE_WATCH_INTERVAL_SECONDS = 10 INVOICE_WATCH_BATCH_LIMIT = 100 -async def periodic_invoice_watcher() -> None: - """Background task: detect paid Lightning invoices and credit balances. +async def _process_invoice_watch_batch(session: AsyncSession) -> None: + result = await session.exec( + select(LightningInvoice) + .where( + col(LightningInvoice.status).in_(_RETRYABLE_INVOICE_STATUSES) + ) + .limit(INVOICE_WATCH_BATCH_LIMIT) + ) + for invoice in result.all(): + try: + definitively_unpaid = await check_invoice_payment(invoice, session) + await _expire_invoice_if_authoritatively_unpaid( + invoice, session, definitively_unpaid + ) + except Exception as e: + logger.error( + "Invoice watcher failed for invoice", + extra={"invoice_id": invoice.id, "error": str(e)}, + ) - Removes the need for clients to poll the status endpoint after paying. - """ + +async def periodic_invoice_watcher() -> None: + """Background task: detect paid Lightning invoices and credit balances.""" while True: try: async with create_session() as session: - now = int(time.time()) - result = await session.exec( - select(LightningInvoice) - .where( - LightningInvoice.status == "pending", - col(LightningInvoice.expires_at) > now, - ) - .limit(INVOICE_WATCH_BATCH_LIMIT) - ) - pending = result.all() - for invoice in pending: - try: - await check_invoice_payment(invoice, session) - except Exception as e: - logger.error( - "Invoice watcher failed for invoice", - extra={"invoice_id": invoice.id, "error": str(e)}, - ) + await _process_invoice_watch_batch(session) except asyncio.CancelledError: raise except Exception as e: diff --git a/routstr/mint.py b/routstr/mint.py index a676a320..c0f0ea3c 100644 --- a/routstr/mint.py +++ b/routstr/mint.py @@ -72,7 +72,15 @@ class MintRateGuard: concurrency = settings.mint_max_concurrency guard = cls._guards.get(mint_url) if guard is None or guard._max_concurrency != concurrency: + previous = guard guard = cls(mint_url, concurrency) + if previous is not None: + # Concurrency changed at runtime: keep the live cooldown/backoff + # state so an active 429 cooldown is not silently discarded. + guard._cooldown_until = previous._cooldown_until + guard._cooldown_reason = previous._cooldown_reason + guard._consecutive_rate_limits = previous._consecutive_rate_limits + guard._needs_probe = previous._needs_probe cls._guards[mint_url] = guard return guard @@ -124,10 +132,9 @@ class MintRateGuard: return self._cooldown_reason if self.cooldown_remaining() > 0 else None def _raise_if_wait_forbidden(self) -> None: - if _fail_fast_depth.get() and ( - self._needs_probe or self.cooldown_remaining() > 0 - ): - raise MintCooldownError(self._mint_url, self.cooldown_remaining()) + remaining = self.cooldown_remaining() + if _fail_fast_depth.get() and remaining > 0: + raise MintCooldownError(self._mint_url, remaining) async def _wait_for_cooldown(self) -> None: while True: @@ -157,11 +164,7 @@ class MintRateGuard: retry_after = None if isinstance(error, httpx.HTTPStatusError): retry_after = parse_retry_after(error.response.headers) - delay = max( - _MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS, - retry_after or 0.0, - ) - self.apply_cooldown(delay, reason="rate_limited") + self.apply_rate_limit_cooldown(retry_after) else: self.apply_cooldown(1.0) logger.warning( @@ -194,6 +197,8 @@ class MintRateGuard: while True: self._raise_if_wait_forbidden() if self._needs_probe or self.cooldown_remaining() > 0: + if _fail_fast_depth.get() and self._probe_lock.locked(): + raise MintCooldownError(self._mint_url, self.cooldown_remaining()) async with self._probe_lock: self._raise_if_wait_forbidden() if self.cooldown_remaining() > 0: diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index a3ab1fb0..702ab527 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -242,7 +242,7 @@ async def calculate_discounted_max_cost( }, ) - return max(0, adjusted) + return max(settings.min_request_msat, adjusted) def estimate_tokens(messages: list) -> int: diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index c6c3ccb7..bb2e6313 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -7,7 +7,11 @@ import httpx from cashu.core.base import MeltQuoteState from cashu.wallet.wallet import Proof, Wallet -from ..mint import MINT_TRANSPORT_EXCEPTIONS, run_mint_operation +from ..mint import ( + MINT_TRANSPORT_EXCEPTIONS, + is_mint_rate_limited, + run_mint_operation, +) try: from bech32 import bech32_decode, convertbits # type: ignore @@ -239,7 +243,15 @@ async def raw_send_to_lnurl( mint_url=str(wallet.url), retry_timeouts=False, ) - except MINT_TRANSPORT_EXCEPTIONS as error: + except Exception as error: + if is_mint_rate_limited(error): + # Cooldown failures happen before dispatch, and HTTP 429 means the + # mint rejected the request. Neither outcome may keep proofs + # reserved as though a Lightning payment could still settle. + await wallet.set_reserved_for_send(proofs, reserved=False) + raise + if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS): + raise melt_response = None melt_error: BaseException | None = error else: diff --git a/routstr/wallet.py b/routstr/wallet.py index c533355e..ef4bf901 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -310,18 +310,10 @@ async def _redeem_same_mint( op_name="redeem_load_mint", mint_url=token_obj.mint, ) - wallet.verify_proofs_dleq(token_obj.proofs) - input_fees = wallet.get_fees_for_proofs(token_obj.proofs) - await run_mint_operation( - lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), - op_name="redeem_split", - mint_url=token_obj.mint, - retry_timeouts=False, - ) except Exception as error: if is_mint_connection_error(error): logger.warning( - "Same-mint redemption failed; client must use a different token", + "Same-mint redemption failed before swap dispatch", extra={ "event": "cashu_same_mint_redemption_failed", "source_mint": token_obj.mint, @@ -338,23 +330,79 @@ async def _redeem_same_mint( ) from error raise + wallet.verify_proofs_dleq(token_obj.proofs) + input_fees = wallet.get_fees_for_proofs(token_obj.proofs) + try: + await run_mint_operation( + lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True), + op_name="redeem_split", + mint_url=token_obj.mint, + retry_timeouts=False, + ) + except Exception as error: + if isinstance(error, httpx.ConnectError): + raise SourceMintConnectionError( + "Issuing Cashu mint is unreachable" + ) from error + if is_mint_connection_error(error): + logger.critical( + "Same-mint swap outcome is ambiguous; sealing source token", + extra={ + "event": "cashu_same_mint_redemption_ambiguous", + "source_mint": token_obj.mint, + "source_unit": token_obj.unit, + "source_amount": token_obj.amount, + "action": "manual_reconciliation_required", + "error": str(error), + "error_type": type(error).__name__, + }, + ) + raise TokenConsumedError( + "Same-mint swap outcome is ambiguous; reconciliation required" + ) from error + raise + return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint async def recieve_token( token: str, + destination_mint: str | None = None, + destination_unit: str | None = None, ) -> tuple[int, str, str]: # amount, unit, mint_url + """Redeem a token while serializing all wallet proof mutation.""" + async with wallet_operation_guard(): + return await _recieve_token_locked(token, destination_mint, destination_unit) + + +async def _recieve_token_locked( + token: str, + destination_mint: str | None = None, + destination_unit: str | None = None, +) -> tuple[int, str, str]: token_obj = deserialize_token_from_string(token) if len(token_obj.keysets) > 1: raise ValueError("Multiple keysets per token currently not supported") + destinations = ( + [destination_mint] + if destination_mint is not None + else list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) + ) + output_unit = ( + token_obj.unit + if token_obj.mint in destinations + else settings.primary_mint_unit + ) + if destination_unit is not None and output_unit != destination_unit: + raise ValueError( + "Cashu token unit does not match the API key liability unit: " + f"expected {destination_unit}, got {output_unit}" + ) + wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) wallet.keyset_id = token_obj.keysets[0] - - if token_obj.mint not in settings.cashu_mints: - destinations = list( - dict.fromkeys([settings.primary_mint, *settings.cashu_mints]) - ) + if token_obj.mint not in destinations: logger.info( "Cashu cross-mint swap required", extra={ @@ -365,7 +413,9 @@ async def recieve_token( "destination_candidates": destinations, }, ) - return await swap_to_trusted_mint(token_obj, wallet) + return await swap_to_trusted_mint( + token_obj, wallet, destination_mints=destinations + ) logger.info( "Trying same-mint Cashu redemption", @@ -382,7 +432,16 @@ async def recieve_token( async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: """Create a token from the preferred mint or another funded trusted mint.""" - effective_mint_url = await find_trusted_mint_with_funds(amount, unit, mint_url) + async with wallet_operation_guard(): + return await _send_locked(amount, unit, mint_url) + + +async def _send_locked( + amount: int, unit: str, mint_url: str | None = None +) -> tuple[int, str]: + effective_mint_url = await find_trusted_mint_with_funds( + amount, unit, mint_url, force_reload=True + ) wallet = await get_wallet(effective_mint_url, unit) proofs = get_proofs_per_mint_and_unit( wallet, effective_mint_url, unit, not_reserved=True @@ -436,16 +495,20 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str async def release_token_reservation(token: str) -> None: """Release a token that was created locally but never handed off.""" - token_obj = deserialize_token_from_string(token) - wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) - await wallet.set_reserved_for_send(token_obj.proofs, reserved=False) + async with wallet_operation_guard(): + token_obj = deserialize_token_from_string(token) + wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) + # This is a local wallet-DB refresh; reservation release must still work + # while the mint is unavailable or cooling down. + await wallet.load_proofs(reload=True) + await wallet.set_reserved_for_send(token_obj.proofs, reserved=False) - secrets = {proof.secret for proof in token_obj.proofs} - for proof in token_obj.proofs: - proof.reserved = False - for proof in wallet.proofs: - if proof.secret in secrets: + secrets = {proof.secret for proof in token_obj.proofs} + for proof in token_obj.proofs: proof.reserved = False + for proof in wallet.proofs: + if proof.secret in secrets: + proof.reserved = False def token_mint_url(token: str, fallback: str | None = None) -> str: @@ -458,7 +521,11 @@ def token_mint_url(token: str, fallback: str | None = None) -> str: async def find_trusted_mint_with_funds( - amount: int, unit: str, preferred_mint: str | None = None + amount: int, + unit: str, + preferred_mint: str | None = None, + *, + force_reload: bool = False, ) -> str: """Choose a trusted mint that can cover a refund without waiting on cooldown.""" trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) @@ -472,7 +539,12 @@ async def find_trusted_mint_with_funds( if mint_cooldown_remaining(mint_url) > 0: continue try: - wallet = await get_wallet(mint_url, unit, retry_on_rate_limit=False) + wallet = await get_wallet( + mint_url, + unit, + retry_on_rate_limit=False, + force_reload=force_reload, + ) except Exception as error: if is_mint_connection_error(error) or is_mint_rate_limited(error): balances[mint_url] = 0 @@ -566,8 +638,27 @@ def _melt_insufficient_shortfall(error: Exception) -> int | None: return 1 +def _trusted_destination_candidates( + candidates: list[str] | None = None, +) -> list[str]: + trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) + if candidates is None: + return trusted + selected = list(dict.fromkeys(candidates)) + untrusted = [mint_url for mint_url in selected if mint_url not in trusted] + if untrusted: + raise ValueError(f"Untrusted destination mint: {untrusted[0]}") + if not selected: + raise ValueError("At least one trusted destination mint is required") + return selected + + async def _request_mint_with_fallback( - amount: int, *, op_name: str, primary_wallet: Wallet | None = None + amount: int, + *, + op_name: str, + primary_wallet: Wallet | None = None, + destination_mints: list[str] | None = None, ) -> tuple[Wallet, str, MintQuote]: """Try request_mint on the primary mint, fall back to other trusted mints on transport or rate-limit failure. Returns the wallet, mint_url, and quote. @@ -581,7 +672,7 @@ async def _request_mint_with_fallback( f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " f"Token value is too small after fee deduction or unit conversion." ) - candidates = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) + candidates = _trusted_destination_candidates(destination_mints) logger.info( "Trying trusted destination mints", extra={ @@ -692,6 +783,7 @@ async def _calculate_swap_amount( token_wallet: Wallet, primary_wallet: Wallet | None, proofs: list, + destination_mints: list[str] | None = None, ) -> int: """ Calculate the amount to mint on the primary mint after accounting for @@ -749,6 +841,7 @@ async def _calculate_swap_amount( receive_amount, op_name="swap_fee_est_mint_quote", primary_wallet=primary_wallet, + destination_mints=destination_mints, ) stage = "source_fee_quote" dummy_melt_quote = await run_mint_operation( @@ -869,7 +962,10 @@ async def _confirm_melt_paid( async def swap_to_trusted_mint( - token_obj: Token, token_wallet: Wallet + token_obj: Token, + token_wallet: Wallet, + *, + destination_mints: list[str] | None = None, ) -> tuple[int, str, str]: logger.info( "Starting Cashu cross-mint swap", @@ -893,10 +989,11 @@ async def swap_to_trusted_mint( amount_msat = token_amount else: raise ValueError("Invalid unit") - # If the token is already from the primary mint, we don't need a cross-mint - # swap — redeem it same-mint. There's no melt/Lightning fee, but the mint's - # NUT-02 input fee still applies; _redeem_same_mint accounts for it. - if token_obj.mint == settings.primary_mint: + destination_candidates = _trusted_destination_candidates(destination_mints) + # If the token is already from an allowed destination, redeem it same-mint. + # There's no melt/Lightning fee, but the mint's NUT-02 input fee still + # applies; _redeem_same_mint accounts for it. + if token_obj.mint in destination_candidates: logger.info( "swap_to_trusted_mint: token already on primary mint, skipping swap", extra={ @@ -916,6 +1013,7 @@ async def swap_to_trusted_mint( token_wallet, primary_wallet, token_obj.proofs, + destination_candidates, ) # The estimate above is non-binding: the mint may demand a higher fee on the @@ -949,6 +1047,7 @@ async def swap_to_trusted_mint( minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet, + destination_mints=destination_candidates, ) logger.info( "swap_to_trusted_mint: mint quote received", @@ -1245,7 +1344,14 @@ async def _credit_balance_locked( ) try: - amount, unit, mint_url = await recieve_token(cashu_token) + destination_mint = key.refund_mint_url or settings.primary_mint + amount, unit, mint_url = await recieve_token( + cashu_token, + destination_mint=destination_mint, + destination_unit=key.refund_currency + if isinstance(key.refund_currency, str) + else None, + ) original_amount = amount original_unit = unit logger.info( @@ -1284,10 +1390,19 @@ async def _credit_balance_locked( # retryable/token-error taxonomy. try: # Atomic UPDATE to prevent race conditions during concurrent topups. + updates: dict[str, object] = { + "balance": db.ApiKey.balance + amount, + } + # Legacy keys may predate refund provenance. Pin them to the + # destination used for this credit before exposing the balance. + if key.refund_mint_url is None: + updates["refund_mint_url"] = mint_url + if key.refund_currency is None: + updates["refund_currency"] = unit stmt = ( update(db.ApiKey) .where(col(db.ApiKey.hashed_key) == key.hashed_key) - .values(balance=(db.ApiKey.balance) + amount) + .values(**updates) ) result = await session.exec(stmt) # type: ignore[call-overload] # If pruning removed this key after redemption, do not commit a no-op @@ -1355,6 +1470,7 @@ async def get_wallet( unit: str = "sat", load: bool = True, retry_on_rate_limit: bool = True, + force_reload: bool = False, ) -> Wallet: global _wallets, _wallet_last_load, _wallet_load_locks id = f"{mint_url}_{unit}" @@ -1366,7 +1482,11 @@ async def get_wallet( if load: now = time.monotonic() last = _wallet_last_load.get(id) - if last is None or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: + if ( + force_reload + or last is None + or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS + ): await run_mint_operation( lambda: _wallets[id].load_mint(), op_name="load_mint", @@ -1849,7 +1969,8 @@ async def _refund_sweep_once(cutoff: int) -> None: claim_owned = col(db.CashuTransaction.sweep_started_at) == claim_started_at redeemed = False try: - await recieve_token(refund.token) + async with wallet_operation_guard(): + await recieve_token(refund.token) redeemed = True finalized = await _set_refund_sweep_state( refund.id, @@ -1980,66 +2101,73 @@ async def periodic_routstr_fee_payout() -> None: continue paid_msats = _sats_to_msats(accumulated_sats) - # Wallet/proof preparation cannot send funds, so do it before the - # durable checkpoint. A preparation failure must not strand an - # in-progress payout that requires manual reconciliation. - wallet = await get_wallet(settings.primary_mint, "sat") - proofs = get_proofs_per_mint_and_unit( - wallet, settings.primary_mint, "sat", not_reserved=True - ) - - async with db.create_session() as session: - payout_checkpointed = await db.reset_routstr_fee(session, paid_msats) - if not payout_checkpointed: - logger.warning("Routstr fee payout was already claimed") - continue - - try: - amount_received = await raw_send_to_lnurl( - wallet, - proofs, - ROUTSTR_LN_ADDRESS, - "sat", - amount=accumulated_sats, + # Serialize proof refresh, reservation, sending, and checkpoint + # finalization with every other wallet mutation across workers. + async with wallet_operation_guard(): + # Wallet/proof preparation cannot send funds, so do it before + # the durable checkpoint. Force a DB reload after taking the + # guard so another worker's reservations are visible. + wallet = await get_wallet( + settings.primary_mint, "sat", force_reload=True ) - except BaseException as e: - logger.critical( - "Routstr fee payout outcome is unknown; manual reconciliation required", - extra={"payout_in_progress_msats": paid_msats}, - exc_info=isinstance(e, Exception), + proofs = get_proofs_per_mint_and_unit( + wallet, settings.primary_mint, "sat", not_reserved=True ) - if not isinstance(e, Exception): - raise - continue - try: async with db.create_session() as session: - payout_completed = await db.complete_routstr_fee_payout( + payout_checkpointed = await db.reset_routstr_fee( session, paid_msats ) - except BaseException as e: - logger.critical( - "Routstr fee payout sent but checkpoint was not completed", - extra={"payout_in_progress_msats": paid_msats}, - exc_info=isinstance(e, Exception), - ) - if not isinstance(e, Exception): - raise - continue - if not payout_completed: - logger.critical( - "Routstr fee payout sent but checkpoint was not completed", - extra={"payout_in_progress_msats": paid_msats}, - ) - continue + if not payout_checkpointed: + logger.warning("Routstr fee payout was already claimed") + continue - logger.info( - "Routstr fee payout sent", - extra={ - "accumulated_sats": accumulated_sats, - "amount_received": amount_received, - }, - ) + try: + amount_received = await raw_send_to_lnurl( + wallet, + proofs, + ROUTSTR_LN_ADDRESS, + "sat", + amount=accumulated_sats, + ) + except BaseException as e: + logger.critical( + "Routstr fee payout outcome is unknown; manual reconciliation required", + extra={"payout_in_progress_msats": paid_msats}, + exc_info=isinstance(e, Exception), + ) + if not isinstance(e, Exception): + raise + continue + + try: + async with db.create_session() as session: + payout_completed = await db.complete_routstr_fee_payout( + session, paid_msats + ) + except BaseException as e: + logger.critical( + "Routstr fee payout sent but checkpoint was not completed", + extra={"payout_in_progress_msats": paid_msats}, + exc_info=isinstance(e, Exception), + ) + if not isinstance(e, Exception): + raise + continue + if not payout_completed: + logger.critical( + "Routstr fee payout sent but checkpoint was not completed", + extra={"payout_in_progress_msats": paid_msats}, + ) + continue + + logger.info( + "Routstr fee payout sent", + extra={ + "accumulated_sats": accumulated_sats, + "amount_received": amount_received, + }, + ) except Exception as e: logger.error( f"Error in Routstr fee payout: {type(e).__name__}", @@ -2048,11 +2176,18 @@ async def periodic_routstr_fee_payout() -> None: async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: - mint = await find_trusted_mint_with_funds(amount, unit, mint) - wallet = await get_wallet(mint, unit) - available = get_proofs_per_mint_and_unit(wallet, mint, unit, not_reserved=True) - proofs, _ = await wallet.select_to_send(available, amount, set_reserved=True) - return await raw_send_to_lnurl(wallet, proofs, address, unit) + async with wallet_operation_guard(): + mint = await find_trusted_mint_with_funds( + amount, unit, mint, force_reload=True + ) + wallet = await get_wallet(mint, unit) + available = get_proofs_per_mint_and_unit( + wallet, mint, unit, not_reserved=True + ) + proofs, _ = await wallet.select_to_send( + available, amount, set_reserved=True + ) + return await raw_send_to_lnurl(wallet, proofs, address, unit) # class Payment: diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index b9d1c25c..aa10a81c 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -203,8 +203,13 @@ class TestmintWallet: token_base64 = base64.urlsafe_b64encode(token_json.encode()).decode() return f"cashuA{token_base64}" - async def redeem_token(self, token: str) -> Tuple[int, str, str]: - """Redeem a Cashu token - compatible with wallet.recieve_token""" + async def redeem_token( + self, + token: str, + destination_mint: str | None = None, + destination_unit: str | None = None, + ) -> Tuple[int, str, str]: + """Redeem a Cashu token - compatible with wallet.recieve_token.""" if not self.wallet: await self.init() diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index 92ee94e7..1a6d94b9 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -246,7 +246,7 @@ async def test_concurrent_payment_checks_mint_and_credit_invoice_once( @pytest.mark.asyncio -async def test_failed_mint_keeps_invoice_pending_for_retry( +async def test_failed_mint_marks_invoice_for_settlement_retry( integration_engine: AsyncEngine, patched_db_engine: None, ) -> None: @@ -270,7 +270,7 @@ async def test_failed_mint_keeps_invoice_pending_for_retry( async with AsyncSession(integration_engine, expire_on_commit=False) as verify: stored = await verify.get(LightningInvoice, invoice.id) assert stored is not None - assert stored.status == "pending" + assert stored.status == "settlement_pending" @pytest.mark.asyncio @@ -399,14 +399,14 @@ async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation( assert sibling_state is not None assert stored_state.expired is False assert sibling_state.expired is False - assert stored.status == "pending" + assert stored.status == "settlement_pending" assert stored_sibling.id == sibling.id assert wallet.mint.await_count == 1 async with AsyncSession(integration_engine, expire_on_commit=False) as verify: stored = await verify.get(LightningInvoice, invoice.id) assert stored is not None - assert stored.status == "pending" + assert stored.status == "settlement_pending" @pytest.mark.asyncio diff --git a/tests/integration/test_lightning_settlement.py b/tests/integration/test_lightning_settlement.py index 4dd7dbc4..f38d6623 100644 --- a/tests/integration/test_lightning_settlement.py +++ b/tests/integration/test_lightning_settlement.py @@ -11,6 +11,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.db import ApiKey, LightningInvoice from routstr.lightning import ( + _expire_invoice_if_authoritatively_unpaid, _finalize_invoice_settlement, _InvoiceSettlement, check_invoice_payment, @@ -251,7 +252,7 @@ async def test_check_invoice_payment_retries_after_mint_success_and_db_failure( pending = await verify.get(LightningInvoice, invoice.id) unchanged = await verify.get(ApiKey, key_hash) assert pending is not None - assert pending.status == "pending" + assert pending.status == "settlement_pending" assert unchanged is not None assert unchanged.balance == 100_000 @@ -271,3 +272,94 @@ async def test_check_invoice_payment_retries_after_mint_success_and_db_failure( wallet.mint.assert_awaited_once_with(100, quote_id=invoice.payment_hash) wallet.restore_tokens_for_keyset.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_expiry_cas_cannot_overwrite_concurrent_paid_invoice( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _lightning_invoice(expires_at=0) + async with AsyncSession(integration_engine, expire_on_commit=False) as seed: + seed.add(invoice) + await seed.commit() + + async with AsyncSession(integration_engine, expire_on_commit=False) as caller: + stale = await caller.get(LightningInvoice, invoice.id) + assert stale is not None + await caller.commit() + + async with AsyncSession(integration_engine, expire_on_commit=False) as paid: + result = await paid.exec( # type: ignore[call-overload] + update(LightningInvoice) + .where(col(LightningInvoice.id) == invoice.id) + .values(status="paid", paid_at=123) + ) + assert result.rowcount == 1 + await paid.commit() + + expired = await _expire_invoice_if_authoritatively_unpaid( + stale, caller, True + ) + + assert expired is False + assert stale.status == "paid" + assert stale.paid_at == 123 + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + stored = await verify.get(LightningInvoice, invoice.id) + assert stored is not None + assert stored.status == "paid" + assert stored.paid_at == 123 + + +@pytest.mark.asyncio +async def test_paid_quote_worker_does_not_mint_after_expiry_claim_wins( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _lightning_invoice(expires_at=0) + async with AsyncSession(integration_engine, expire_on_commit=False) as seed: + seed.add(invoice) + await seed.commit() + + quote_started = asyncio.Event() + release_quote = asyncio.Event() + + async def paid_quote_after_expiry(*_args: object, **_kwargs: object) -> Mock: + quote_started.set() + await release_quote.wait() + return Mock(paid=True) + + wallet = Mock( + get_mint_quote=AsyncMock(side_effect=paid_quote_after_expiry), + mint=AsyncMock(), + ) + + async with AsyncSession(integration_engine, expire_on_commit=False) as worker: + observed_pending = await worker.get(LightningInvoice, invoice.id) + assert observed_pending is not None + + with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)): + settlement_task = asyncio.create_task( + check_invoice_payment(observed_pending, worker) + ) + await quote_started.wait() + + async with AsyncSession( + integration_engine, expire_on_commit=False + ) as expirer: + expiry_view = await expirer.get(LightningInvoice, invoice.id) + assert expiry_view is not None + await expirer.commit() + assert await _expire_invoice_if_authoritatively_unpaid( + expiry_view, expirer, True + ) + + release_quote.set() + assert await settlement_task is False + + wallet.mint.assert_not_awaited() + async with AsyncSession(integration_engine, expire_on_commit=False) as verify: + stored = await verify.get(LightningInvoice, invoice.id) + assert stored is not None + assert stored.status == "expired" diff --git a/tests/integration/test_periodic_payout_safety.py b/tests/integration/test_periodic_payout_safety.py index 0f2769fc..62e5e55f 100644 --- a/tests/integration/test_periodic_payout_safety.py +++ b/tests/integration/test_periodic_payout_safety.py @@ -110,7 +110,11 @@ async def test_payout_does_not_send_proofs_whose_liability_commit_is_in_flight( finish_redemption = asyncio.Event() liability_read = asyncio.Event() - async def redeem_token(token: str) -> tuple[int, str, str]: + async def redeem_token( + token: str, + destination_mint: str | None = None, + destination_unit: str | None = None, + ) -> tuple[int, str, str]: proofs.append(MagicMock(amount=200)) proof_visible.set() await finish_redemption.wait() diff --git a/tests/integration/test_prune_dead_api_keys.py b/tests/integration/test_prune_dead_api_keys.py index 4aa95175..85bb13b0 100644 --- a/tests/integration/test_prune_dead_api_keys.py +++ b/tests/integration/test_prune_dead_api_keys.py @@ -126,8 +126,11 @@ async def test_parent_and_child_keys_are_not_pruned( @pytest.mark.asyncio -async def test_pending_invoice_protects_key(patched_db_engine: None) -> None: - """A key referenced by a pending topup invoice is never pruned mid-topup.""" +@pytest.mark.parametrize("status", ["pending", "settlement_pending"]) +async def test_retryable_invoice_protects_key( + patched_db_engine: None, status: str +) -> None: + """A key referenced by a retryable topup invoice is never pruned mid-topup.""" key = _dead_key(LONG_AGO) invoice = LightningInvoice( id=f"inv_{uuid.uuid4().hex}", @@ -135,7 +138,7 @@ async def test_pending_invoice_protects_key(patched_db_engine: None) -> None: amount_sats=10, description="topup", payment_hash=uuid.uuid4().hex, - status="pending", + status=status, api_key_hash=key.hashed_key, purpose="topup", expires_at=NOW + 10_000, diff --git a/tests/integration/test_swap_fee_retry.py b/tests/integration/test_swap_fee_retry.py index b27360a2..8195a7a8 100644 --- a/tests/integration/test_swap_fee_retry.py +++ b/tests/integration/test_swap_fee_retry.py @@ -29,7 +29,9 @@ from routstr.core.settings import settings # with the testmint stub that bypasses swapping (see conftest.py). from routstr.wallet import recieve_token as _real_recieve_token -PRIMARY_MINT = "http://primary:3338" +# Match the authenticated fixture's persisted refund mint: existing-key topups +# are intentionally constrained to that mint for collateral provenance. +PRIMARY_MINT = "http://localhost:3338" def _make_swap_mocks( diff --git a/tests/unit/test_admin_withdraw.py b/tests/unit/test_admin_withdraw.py index e54f7d01..07a98516 100644 --- a/tests/unit/test_admin_withdraw.py +++ b/tests/unit/test_admin_withdraw.py @@ -1,8 +1,12 @@ +import base64 +import json from types import SimpleNamespace from unittest.mock import AsyncMock, Mock import pytest +from fastapi import HTTPException +import routstr.wallet as wallet_module from routstr.core import admin @@ -13,20 +17,12 @@ async def test_withdraw_uses_effective_mint_and_records_outgoing_transaction( ) -> None: primary_mint = "https://primary.example" effective_mint = requested_mint or primary_mint - wallet = object() - proofs = [SimpleNamespace(amount=40), SimpleNamespace(amount=60)] token = "cashuBoutgoing" - - get_wallet = AsyncMock(return_value=wallet) - get_proofs = Mock(return_value=proofs) - filter_proofs = AsyncMock(return_value=proofs) send_token = AsyncMock(return_value=token) store_transaction = AsyncMock(return_value=True) - monkeypatch.setattr(admin, "get_wallet", get_wallet) - monkeypatch.setattr(admin, "get_proofs_per_mint_and_unit", get_proofs) - monkeypatch.setattr(admin, "slow_filter_spend_proofs", filter_proofs) monkeypatch.setattr(admin, "send_token", send_token) + monkeypatch.setattr(admin, "token_mint_url", Mock(return_value=effective_mint)) monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction) monkeypatch.setattr(admin.settings, "primary_mint", primary_mint) @@ -35,10 +31,7 @@ async def test_withdraw_uses_effective_mint_and_records_outgoing_transaction( admin.WithdrawRequest(amount=75, mint_url=requested_mint, unit="sat"), ) - assert result == {"token": token} - get_wallet.assert_awaited_once_with(effective_mint, "sat") - get_proofs.assert_called_once_with(wallet, effective_mint, "sat", not_reserved=True) - filter_proofs.assert_awaited_once_with(proofs, wallet) + assert result == {"token": token, "mint_url": effective_mint} send_token.assert_awaited_once_with(75, "sat", effective_mint) store_transaction.assert_awaited_once_with( token=token, @@ -56,17 +49,10 @@ async def test_withdraw_returns_issued_token_when_audit_storage_fails( monkeypatch: pytest.MonkeyPatch, ) -> None: mint = "https://primary.example" - proofs = [SimpleNamespace(amount=100)] token = "cashuBrecoverable" - monkeypatch.setattr(admin, "get_wallet", AsyncMock(return_value=object())) - monkeypatch.setattr( - admin, "get_proofs_per_mint_and_unit", Mock(return_value=proofs) - ) - monkeypatch.setattr( - admin, "slow_filter_spend_proofs", AsyncMock(return_value=proofs) - ) monkeypatch.setattr(admin, "send_token", AsyncMock(return_value=token)) + monkeypatch.setattr(admin, "token_mint_url", Mock(return_value=mint)) monkeypatch.setattr( admin, "store_cashu_transaction", @@ -78,5 +64,89 @@ async def test_withdraw_returns_issued_token_when_audit_storage_fails( result = await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75)) - assert result == {"token": token} + assert result == {"token": token, "mint_url": mint} critical.assert_called_once() + + +@pytest.mark.asyncio +async def test_withdraw_falls_back_from_insufficient_preferred_mint( + monkeypatch: pytest.MonkeyPatch, +) -> None: + requested_mint = "https://primary.example" + actual_mint = "https://secondary.example" + proofs = [SimpleNamespace(amount=100, reserved=False, id="00")] + token_payload = { + "token": [ + { + "mint": actual_mint, + "proofs": [ + { + "id": "00", + "amount": 75, + "secret": "secret", + "C": "02" + "00" * 32, + } + ], + } + ], + "unit": "sat", + } + token = "cashuA" + base64.urlsafe_b64encode( + json.dumps(token_payload).encode() + ).decode() + wallet = SimpleNamespace( + keysets={}, + proofs=proofs, + select_to_send=AsyncMock(return_value=(proofs, 0)), + serialize_proofs=AsyncMock(return_value=token), + set_reserved_for_send=AsyncMock(), + ) + find_funded = AsyncMock(return_value=actual_mint) + store_transaction = AsyncMock(return_value=True) + + monkeypatch.setattr(wallet_module, "find_trusted_mint_with_funds", find_funded) + monkeypatch.setattr(wallet_module, "get_wallet", AsyncMock(return_value=wallet)) + monkeypatch.setattr( + wallet_module, "get_proofs_per_mint_and_unit", Mock(return_value=proofs) + ) + monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction) + + result = await admin.withdraw( + Mock(), admin.WithdrawRequest(amount=75, mint_url=requested_mint) + ) + + assert result == {"token": token, "mint_url": actual_mint} + find_funded.assert_awaited_once_with( + 75, "sat", requested_mint, force_reload=True + ) + wallet.select_to_send.assert_awaited_once() + store_transaction.assert_awaited_once_with( + token=token, + amount=75, + unit="sat", + mint_url=actual_mint, + typ="out", + collected=False, + source="admin", + ) + + +@pytest.mark.asyncio +async def test_withdraw_maps_true_aggregate_insufficient_funds_to_400( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + admin, + "send_token", + AsyncMock( + side_effect=ValueError( + "No trusted mint has 75 sat available; balances={'mint': 0}" + ) + ), + ) + + with pytest.raises(HTTPException) as exc_info: + await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75)) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail == "Insufficient wallet balance" diff --git a/tests/unit/test_auth_cashu.py b/tests/unit/test_auth_cashu.py index 1d3e6396..29421231 100644 --- a/tests/unit/test_auth_cashu.py +++ b/tests/unit/test_auth_cashu.py @@ -270,6 +270,54 @@ async def test_internal_error_with_invalid_keyword_does_not_masquerade( assert await session.get(ApiKey, hashed_key) is None +@pytest.mark.asyncio +async def test_primary_msat_token_sets_provenance_without_cashu_mint_duplicate( + session: AsyncSession, +) -> None: + token = "cashuAprimary_msat_token" + token_obj = SimpleNamespace(mint="http://primary:3338", unit="msat") + credit = AsyncMock(return_value=1_000) + + from routstr.core.settings import settings + + with ( + patch.object(settings, "primary_mint", token_obj.mint), + patch.object(settings, "primary_mint_unit", "msat"), + patch.object(settings, "cashu_mints", []), + patch("routstr.auth.deserialize_token_from_string", return_value=token_obj), + patch("routstr.auth.credit_balance", new=credit), + ): + key = await validate_bearer_key(token, session) + + assert key.refund_mint_url == token_obj.mint + assert key.refund_currency == "msat" + credit.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_primary_token_unit_mismatch_is_rejected_before_redemption( + session: AsyncSession, +) -> None: + token = "cashuAprimary_wrong_unit" + token_obj = SimpleNamespace(mint="http://primary:3338", unit="sat") + credit = AsyncMock(return_value=1_000) + + from routstr.core.settings import settings + + with ( + patch.object(settings, "primary_mint", token_obj.mint), + patch.object(settings, "primary_mint_unit", "msat"), + patch.object(settings, "cashu_mints", []), + patch("routstr.auth.deserialize_token_from_string", return_value=token_obj), + patch("routstr.auth.credit_balance", new=credit), + ): + with pytest.raises(HTTPException) as exc_info: + await validate_bearer_key(token, session) + + assert exc_info.value.status_code == 400 + credit.assert_not_awaited() + + @pytest.mark.asyncio async def test_malformed_cashu_token_returns_400_invalid_token( session: AsyncSession, diff --git a/tests/unit/test_coverage_admin.py b/tests/unit/test_coverage_admin.py index c198ceae..67e04942 100644 --- a/tests/unit/test_coverage_admin.py +++ b/tests/unit/test_coverage_admin.py @@ -4,7 +4,7 @@ Tests admin endpoints that are testable without full app setup: withdraw validation, authentication guards, and slug validation. """ -from unittest.mock import Mock, patch +from unittest.mock import AsyncMock, patch import pytest from fastapi import HTTPException, Request @@ -46,22 +46,19 @@ async def test_withdraw_rejects_insufficient_balance() -> None: request = Request(scope={"type": "http", "method": "POST"}) - with patch("routstr.core.admin.get_wallet") as mock_wallet, \ - patch("routstr.core.admin.get_proofs_per_mint_and_unit") as mock_proofs, \ - patch("routstr.core.admin.slow_filter_spend_proofs") as mock_filter: - - mock_w = Mock() - mock_w.keysets = {} - mock_w.proofs = [] - mock_wallet.return_value = mock_w - mock_proofs.return_value = [] - mock_filter.return_value = [] - + with patch( + "routstr.core.admin.send_token", + new=AsyncMock( + side_effect=ValueError( + "No trusted mint has 1000000 sat available; balances={}" + ) + ), + ): with pytest.raises(HTTPException) as exc_info: await withdraw(request, WithdrawRequest(amount=1000000, unit="sat")) - assert exc_info.value.status_code == 400 - assert "Insufficient" in str(exc_info.value.detail) + assert exc_info.value.status_code == 400 + assert "Insufficient" in str(exc_info.value.detail) # =========================================================================== diff --git a/tests/unit/test_fee_payout_crash_safety.py b/tests/unit/test_fee_payout_crash_safety.py index ae7d0788..98cbe362 100644 --- a/tests/unit/test_fee_payout_crash_safety.py +++ b/tests/unit/test_fee_payout_crash_safety.py @@ -67,7 +67,7 @@ async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> N payout_wallet = Mock() events: list[str] = [] - async def prepare(*_args: object) -> Mock: + async def prepare(*_args: object, **_kwargs: object) -> Mock: events.append("prepare") return payout_wallet diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py index d8741b7e..584950d6 100644 --- a/tests/unit/test_lightning_settlement.py +++ b/tests/unit/test_lightning_settlement.py @@ -6,13 +6,16 @@ from unittest.mock import AsyncMock, Mock, patch import httpx import pytest -from cashu.core.base import Proof +from cashu.core.base import MintQuoteState, Proof from routstr.lightning import ( + InvoiceRecoverRequest, _invoice_settlement_locks, _is_outputs_already_signed, _mint_invoice_quote, check_invoice_payment, + get_invoice_status, + recover_invoice, ) from routstr.wallet import Wallet @@ -30,6 +33,8 @@ def _invoice(**overrides: object) -> SimpleNamespace: "balance_limit": None, "balance_limit_reset": None, "validity_date": None, + "created_at": 1, + "expires_at": 2, } values.update(overrides) return SimpleNamespace(**values) @@ -159,14 +164,21 @@ async def test_non_pending_invoice_is_not_minted() -> None: @pytest.mark.asyncio -async def test_ambiguous_invoice_mint_timeout_does_not_expose_paid() -> None: +async def test_ambiguous_invoice_mint_timeout_remains_recoverable() -> None: _invoice_settlement_locks.clear() invoice = _invoice() session = AsyncMock() wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True))) + state_session = AsyncMock() + state_session.exec.return_value.rowcount = 1 + + @asynccontextmanager + async def owned_session() -> AsyncIterator[AsyncMock]: + yield state_session with ( patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning.create_session", owned_session), patch( "routstr.lightning._mint_invoice_quote", AsyncMock(side_effect=httpx.TimeoutException("response lost")), @@ -175,12 +187,152 @@ async def test_ambiguous_invoice_mint_timeout_does_not_expose_paid() -> None: ): await check_invoice_payment(invoice, session) # type: ignore[arg-type] - assert invoice.status == "pending" + assert invoice.status == "settlement_pending" + state_session.commit.assert_awaited_once() session.rollback.assert_not_awaited() # One commit closes the initial read transaction before external I/O. session.commit.assert_awaited_once() +@pytest.mark.asyncio +async def test_quote_lookup_timeout_is_not_definitively_unpaid() -> None: + _invoice_settlement_locks.clear() + invoice = _invoice(expires_at=0) + session = AsyncMock() + wallet = Mock( + get_mint_quote=AsyncMock(side_effect=httpx.TimeoutException("quote timeout")) + ) + + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning._reload_invoice_view", AsyncMock()), + ): + result = await check_invoice_payment(invoice, session) # type: ignore[arg-type] + + assert result is False + + +@pytest.mark.asyncio +async def test_overdue_invoice_does_not_expire_after_ambiguous_quote_lookup() -> None: + invoice = _invoice(status="pending", expires_at=0) + session = AsyncMock() + session.get.return_value = invoice + check = AsyncMock(return_value=False) + + with patch("routstr.lightning.check_invoice_payment", check): + response = await get_invoice_status(invoice.id, session) # type: ignore[arg-type] + + assert response.status == "pending" + session.commit.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_overdue_invoice_expires_only_after_definitive_unpaid_quote() -> None: + invoice = _invoice(status="pending", expires_at=0) + session = AsyncMock() + session.get.return_value = invoice + check = AsyncMock(return_value=True) + + async def expire( + candidate: SimpleNamespace, _session: AsyncMock, definitive: bool + ) -> bool: + assert definitive is True + candidate.status = "expired" + return True + + with ( + patch("routstr.lightning.check_invoice_payment", check), + patch( + "routstr.lightning._expire_invoice_if_authoritatively_unpaid", + side_effect=expire, + ) as expire_invoice, + ): + response = await get_invoice_status(invoice.id, session) # type: ignore[arg-type] + + assert response.status == "expired" + expire_invoice.assert_awaited_once_with(invoice, session, True) + session.commit.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_recover_applies_authoritative_expiry_helper() -> None: + invoice = _invoice(status="pending", expires_at=0) + session = AsyncMock() + result = Mock() + result.first.return_value = invoice + session.exec.return_value = result + check = AsyncMock(return_value=True) + + async def expire( + candidate: SimpleNamespace, _session: AsyncMock, definitive: bool + ) -> bool: + assert definitive is True + candidate.status = "expired" + return True + + with ( + patch("routstr.lightning.check_invoice_payment", check), + patch( + "routstr.lightning._expire_invoice_if_authoritatively_unpaid", + side_effect=expire, + ) as expire_invoice, + ): + response = await recover_invoice( + InvoiceRecoverRequest(bolt11="lnbc-test"), session # type: ignore[arg-type] + ) + + assert response.status == "expired" + expire_invoice.assert_awaited_once_with(invoice, session, True) + + +@pytest.mark.asyncio +async def test_paid_state_write_failure_still_reports_non_expirable_outcome() -> None: + _invoice_settlement_locks.clear() + invoice = _invoice(expires_at=0) + session = AsyncMock() + wallet = Mock( + get_mint_quote=AsyncMock( + return_value=Mock(paid=True, state=MintQuoteState.paid) + ) + ) + + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch( + "routstr.lightning._mint_invoice_quote", + AsyncMock(side_effect=httpx.TimeoutException("response lost")), + ), + patch( + "routstr.lightning.create_session", + side_effect=RuntimeError("database unavailable"), + ), + patch("routstr.lightning._reload_invoice_view", AsyncMock()), + ): + definitively_unpaid = await check_invoice_payment( + invoice, session # type: ignore[arg-type] + ) + + assert definitively_unpaid is False + assert invoice.status == "pending" + + +@pytest.mark.asyncio +async def test_settlement_pending_invoice_does_not_expire() -> None: + invoice = _invoice(status="settlement_pending", expires_at=0) + session = AsyncMock() + session.get.return_value = invoice + check = AsyncMock() + + with patch("routstr.lightning.check_invoice_payment", check): + response = await get_invoice_status( + invoice.id, session # type: ignore[arg-type] + ) + + check.assert_awaited_once_with(invoice, session) + assert response.status == "settlement_pending" + session.commit.assert_not_awaited() + + @pytest.mark.asyncio async def test_concurrent_invoice_checks_finalize_once_in_process() -> None: _invoice_settlement_locks.clear() @@ -195,7 +347,9 @@ async def test_concurrent_invoice_checks_finalize_once_in_process() -> None: @asynccontextmanager async def owned_session() -> AsyncIterator[AsyncMock]: - yield AsyncMock() + owned = AsyncMock() + owned.exec.return_value.rowcount = 1 + yield owned with ( patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index 7311b4a2..a1b40bcf 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -4,10 +4,12 @@ import asyncio from typing import Any from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from cashu.core.base import MeltQuoteState from routstr.core.settings import settings +from routstr.mint import MintCooldownError, MintRateGuard from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl LNURL_DATA = { @@ -111,6 +113,76 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None: wallet.get_melt_quote.assert_awaited_once_with("q") +@pytest.mark.asyncio +@pytest.mark.parametrize("rate_error", ["cooldown", "http_429"]) +async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs( + rate_error: str, +) -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock() + wallet.set_reserved_for_send = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + async def run_operation(factory: Any, *, op_name: str, **_: object) -> Any: + if op_name == "lnurl_melt": + if rate_error == "cooldown": + raise MintCooldownError(str(wallet.url), 60) + request = httpx.Request("POST", f"{wallet.url}/v1/melt/bolt11") + response = httpx.Response(429, request=request) + raise httpx.HTTPStatusError( + "rate limited", request=request, response=response + ) + return await factory() + + with ( + data_patch, + invoice_patch, + patch( + "routstr.payment.lnurl.run_mint_operation", + side_effect=run_operation, + ), + pytest.raises((MintCooldownError, httpx.HTTPStatusError)), + ): + await raw_send_to_lnurl( + wallet, proofs, "owner@ln.tld", "sat", amount=1000 + ) + + wallet.melt.assert_not_awaited() + wallet.set_reserved_for_send.assert_awaited_once_with( + proofs, reserved=False + ) + + +@pytest.mark.asyncio +async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None: + wallet, proofs = _wallet() + request = httpx.Request("POST", f"{wallet.url}/v1/melt/bolt11") + response = httpx.Response(429, request=request) + wallet.melt = AsyncMock( + side_effect=httpx.HTTPStatusError( + "rate limited", request=request, response=response + ) + ) + wallet.set_reserved_for_send = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_retry_max_attempts", 0), + data_patch, + invoice_patch, + pytest.raises(httpx.HTTPStatusError), + ): + await raw_send_to_lnurl( + wallet, proofs, "owner@ln.tld", "sat", amount=1000 + ) + + wallet.melt.assert_awaited_once() + wallet.set_reserved_for_send.assert_awaited_once_with( + proofs, reserved=False + ) + MintRateGuard._guards.pop(str(wallet.url), None) + + @pytest.mark.asyncio async def test_raw_send_to_lnurl_succeeds_on_explicit_paid_response() -> None: wallet, proofs = _wallet() diff --git a/tests/unit/test_mint.py b/tests/unit/test_mint.py index 5a258a29..6b70a558 100644 --- a/tests/unit/test_mint.py +++ b/tests/unit/test_mint.py @@ -1,3 +1,4 @@ +import asyncio from unittest.mock import AsyncMock, Mock, patch import httpx @@ -31,6 +32,43 @@ async def test_cooldown_fails_fast_while_wallet_mutation_scope_is_held() -> None sleep.assert_not_awaited() +@pytest.mark.asyncio +async def test_expired_cooldown_allows_probe_in_wallet_mutation_scope() -> None: + guard = MintRateGuard("http://mint:3338", max_concurrency=1) + guard.apply_cooldown(0, reason="rate_limited") + operation = AsyncMock(return_value="recovered") + + async with fail_fast_mint_operations(): + result = await guard.run(operation) + + assert result == "recovered" + operation.assert_awaited_once() + assert guard._needs_probe is False + + +@pytest.mark.asyncio +async def test_fail_fast_does_not_wait_behind_existing_probe() -> None: + guard = MintRateGuard("http://mint:3338", max_concurrency=1) + guard.apply_cooldown(0, reason="rate_limited") + probe_started = asyncio.Event() + release_probe = asyncio.Event() + + async def probe() -> str: + probe_started.set() + await release_probe.wait() + return "recovered" + + first = asyncio.create_task(guard.run(probe)) + await probe_started.wait() + try: + async with fail_fast_mint_operations(): + with pytest.raises(MintCooldownError): + await asyncio.wait_for(guard.run(AsyncMock()), timeout=0.05) + finally: + release_probe.set() + assert await first == "recovered" + + @pytest.mark.asyncio async def test_cashu_429_dispatches_through_wallet_override() -> None: async def handler(request: httpx.Request) -> httpx.Response: @@ -63,3 +101,21 @@ async def test_cashu_429_dispatches_through_wallet_override() -> None: pytest.raises(MintRateLimitedError), ): await wallet.mint_quote(1, Unit.sat) + + +async def test_guard_concurrency_change_preserves_cooldown_state() -> None: + from routstr.core.settings import settings + + mint_url = "https://mint.test-concurrency-carryover" + with patch.object(settings, "mint_max_concurrency", 2): + guard = MintRateGuard.get(mint_url) + guard.apply_cooldown(120.0, reason="rate_limited") + guard._consecutive_rate_limits = 3 + + with patch.object(settings, "mint_max_concurrency", 5): + rebuilt = MintRateGuard.get(mint_url) + + assert rebuilt is not guard + assert rebuilt.cooldown_remaining() > 0 + assert rebuilt._cooldown_reason == "rate_limited" + assert rebuilt._consecutive_rate_limits == 3 diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index ef8dde63..6809d94c 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -125,3 +125,33 @@ async def test_get_max_cost_for_model_tolerance() -> None: "gpt-4", session=mock_session, model_obj=mock_model ) assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000 + + +async def test_discounted_max_cost_floors_at_min_request_msat() -> None: + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.001 + pricing.max_prompt_cost = 100.0 + pricing.max_completion_cost = 100.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = { + "model": "test-model", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 1, + } + + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost = await calculate_discounted_max_cost(150_000, body, model_obj) + + assert cost == 1000 diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 03123a60..e83e306a 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -2,7 +2,8 @@ import asyncio import base64 import json import socket -from collections.abc import Generator +from collections.abc import AsyncIterator, Generator +from contextlib import asynccontextmanager from unittest.mock import AsyncMock, Mock, patch import httpx @@ -61,6 +62,49 @@ async def test_get_balance() -> None: assert balance == 50000 +@pytest.mark.asyncio +async def test_get_wallet_force_reload_bypasses_reload_interval() -> None: + from routstr.wallet import get_wallet + + mock_wallet = Mock(load_mint=AsyncMock(), load_proofs=AsyncMock()) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + await get_wallet("http://mint:3338", "sat") + await get_wallet("http://mint:3338", "sat", force_reload=True) + + assert mock_wallet.load_mint.await_count == 2 + assert mock_wallet.load_proofs.await_count == 2 + + +@pytest.mark.asyncio +async def test_public_recieve_token_holds_wallet_operation_guard() -> None: + inside_guard = False + + @asynccontextmanager + async def operation_guard() -> AsyncIterator[None]: + nonlocal inside_guard + inside_guard = True + try: + yield + finally: + inside_guard = False + + async def receive_locked(*_args: object, **_kwargs: object) -> tuple[int, str, str]: + assert inside_guard + return 1, "sat", "https://mint.example" + + with ( + patch("routstr.wallet.wallet_operation_guard", operation_guard), + patch("routstr.wallet._recieve_token_locked", side_effect=receive_locked), + ): + assert await recieve_token("cashuAtoken") == ( + 1, + "sat", + "https://mint.example", + ) + + assert inside_guard is False + + @pytest.mark.asyncio async def test_recieve_token_valid() -> None: token_data = { @@ -167,6 +211,78 @@ async def test_recieve_token_trusted_mint_deducts_input_fee() -> None: ) +@pytest.mark.asyncio +async def test_recieve_token_uses_only_requested_destination_mint() -> None: + from routstr.core.settings import settings + + source = "http://foreign:3338" + destination = "http://key-mint:3338" + token = Mock( + mint=source, + unit="sat", + amount=100, + keysets=["keyset1"], + proofs=[Mock(amount=100)], + ) + source_wallet = Mock() + swap = AsyncMock(return_value=(99, "sat", destination)) + + with ( + patch.object(settings, "primary_mint", destination), + patch.object(settings, "cashu_mints", [destination]), + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=source_wallet)), + patch("routstr.wallet.swap_to_trusted_mint", swap), + ): + result = await recieve_token( + "cashuAtoken", destination_mint=destination, destination_unit="sat" + ) + + assert result == (99, "sat", destination) + swap.assert_awaited_once_with( + token, source_wallet, destination_mints=[destination] + ) + + +@pytest.mark.asyncio +async def test_recieve_token_rejects_unit_mismatch_before_wallet_mutation() -> None: + token = Mock(mint="http://key-mint:3338", unit="msat", keysets=["keyset"]) + get_wallet = AsyncMock() + + with ( + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch("routstr.wallet.get_wallet", get_wallet), + pytest.raises(ValueError, match="liability unit"), + ): + await recieve_token( + "cashuAtoken", + destination_mint="http://key-mint:3338", + destination_unit="sat", + ) + + get_wallet.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_recieve_token_cross_mint_output_unit_must_match() -> None: + token = Mock(mint="http://foreign:3338", unit="msat", keysets=["keyset"]) + get_wallet = AsyncMock() + + with ( + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch("routstr.wallet.settings.primary_mint_unit", "sat"), + patch("routstr.wallet.get_wallet", get_wallet), + pytest.raises(ValueError, match="liability unit"), + ): + await recieve_token( + "cashuAtoken", + destination_mint="http://key-mint:3338", + destination_unit="msat", + ) + + get_wallet.assert_not_awaited() + + @pytest.mark.asyncio async def test_primary_mint_failure_does_not_try_another_mint() -> None: from routstr.core.settings import settings @@ -207,6 +323,54 @@ async def test_primary_mint_failure_does_not_try_another_mint() -> None: assert failure["action"] == "retry_with_token_from_another_mint" +@pytest.mark.asyncio +async def test_same_mint_split_timeout_is_non_retryable() -> None: + from routstr.wallet import _redeem_same_mint + + token = Mock( + keysets=["keyset1"], + mint="http://mint:3338", + unit="sat", + amount=1000, + proofs=[Mock(amount=1000)], + ) + wallet = Mock( + load_mint=AsyncMock(), + split=AsyncMock(side_effect=httpx.ReadTimeout("response lost")), + get_fees_for_proofs=Mock(return_value=0), + ) + + with pytest.raises(TokenConsumedError, match="outcome is ambiguous") as caught: + await _redeem_same_mint(wallet, token) + + classified = classify_redemption_error(caught.value) + assert classified is not None + assert classified[0] == "token_consumed" + assert classified[1] == 500 + assert classified[3] == "cashu_token_consumed" + + +@pytest.mark.asyncio +async def test_same_mint_split_connect_error_remains_retryable() -> None: + from routstr.wallet import SourceMintConnectionError, _redeem_same_mint + + token = Mock( + keysets=["keyset1"], + mint="http://mint:3338", + unit="sat", + amount=1000, + proofs=[Mock(amount=1000)], + ) + wallet = Mock( + load_mint=AsyncMock(), + split=AsyncMock(side_effect=httpx.ConnectError("connect failed")), + get_fees_for_proofs=Mock(return_value=0), + ) + + with pytest.raises(SourceMintConnectionError): + await _redeem_same_mint(wallet, token) + + @pytest.mark.asyncio async def test_send_token() -> None: mock_wallet = Mock() @@ -224,7 +388,11 @@ async def test_release_token_reservation_unreserves_local_proofs() -> None: token_proof = Mock(secret="proof-secret", reserved=True) cached_proof = Mock(secret="proof-secret", reserved=True) token = Mock(mint="http://mint:3338", unit="sat", proofs=[token_proof]) - wallet = Mock(proofs=[cached_proof], set_reserved_for_send=AsyncMock()) + wallet = Mock( + proofs=[cached_proof], + load_proofs=AsyncMock(), + set_reserved_for_send=AsyncMock(), + ) with ( patch("routstr.wallet.deserialize_token_from_string", return_value=token), patch( @@ -234,6 +402,7 @@ async def test_release_token_reservation_unreserves_local_proofs() -> None: await release_token_reservation("cashu-token") get_wallet.assert_awaited_once_with("http://mint:3338", "sat", load=False) + wallet.load_proofs.assert_awaited_once_with(reload=True) wallet.set_reserved_for_send.assert_awaited_once_with(token.proofs, reserved=False) assert token_proof.reserved is False assert cached_proof.reserved is False @@ -270,6 +439,64 @@ async def test_refund_mint_falls_back_to_trusted_mint_with_funds() -> None: assert mint == secondary +@pytest.mark.asyncio +async def test_send_refreshes_reservations_inside_wallet_guard() -> None: + mint = "http://mint:3338" + proof = Mock(amount=1000, reserved=False) + wallet = Mock( + keysets={}, + proofs=[proof], + select_to_send=AsyncMock(return_value=([proof], None)), + serialize_proofs=AsyncMock(return_value="token"), + set_reserved_for_send=AsyncMock(), + ) + inside_guard = False + + @asynccontextmanager + async def operation_guard() -> AsyncIterator[None]: + nonlocal inside_guard + inside_guard = True + try: + yield + finally: + inside_guard = False + + async def find_mint( + amount: int, + unit: str, + preferred_mint: str | None, + *, + force_reload: bool, + ) -> str: + assert inside_guard + assert (amount, unit, preferred_mint, force_reload) == ( + 1000, + "sat", + mint, + True, + ) + return mint + + async def get_loaded_wallet(*_: object, **__: object) -> Mock: + assert inside_guard + return wallet + + with ( + patch("routstr.wallet.wallet_operation_guard", operation_guard), + patch("routstr.wallet.find_trusted_mint_with_funds", side_effect=find_mint), + patch("routstr.wallet.get_wallet", side_effect=get_loaded_wallet), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + return_value=[proof], + ), + ): + assert await send(1000, "sat", mint) == (1000, "token") + + wallet.set_reserved_for_send.assert_awaited_once_with( + [proof], reserved=True + ) + + @pytest.mark.asyncio async def test_send_falls_back_when_preferred_mint_has_only_reserved_balance() -> None: from routstr.core.settings import settings @@ -395,6 +622,27 @@ async def test_credit_balance() -> None: assert mock_session.refresh.called +@pytest.mark.asyncio +async def test_credit_balance_constrains_redemption_to_key_mint() -> None: + key_mint = "http://key-mint:3338" + mock_key = Mock( + balance=1_000_000, + hashed_key="test_hash", + refund_mint_url=key_mint, + refund_currency="sat", + ) + mock_session = AsyncMock() + mock_session.exec.return_value.rowcount = 1 + receive = AsyncMock(return_value=(1000, "sat", key_mint)) + + with patch("routstr.wallet.recieve_token", receive): + await credit_balance("cashuAtoken", mock_key, mock_session) + + receive.assert_awaited_once_with( + "cashuAtoken", destination_mint=key_mint, destination_unit="sat" + ) + + @pytest.mark.asyncio async def test_credit_balance_rejects_zero_amount() -> None: """A zero/dust redemption must raise BEFORE any commit, so no orphan @@ -2386,12 +2634,12 @@ def test_classify_500_with_rate_limit_text_is_not_mint_rate_limited() -> None: # --------------------------------------------------------------------------- -# _MintRateGuard — probe does NOT escalate cooldown counter +# _MintRateGuard — probe backoff escalation and recovery # --------------------------------------------------------------------------- @pytest.mark.asyncio -async def test_probe_does_not_escalate_consecutive_rate_limits() -> None: +async def test_probe_escalates_consecutive_rate_limits() -> None: from routstr.mint import MintRateGuard guard = MintRateGuard("http://mint", max_concurrency=0) @@ -2401,8 +2649,9 @@ async def test_probe_does_not_escalate_consecutive_rate_limits() -> None: with pytest.raises(httpx.HTTPStatusError): await guard.run(AsyncMock(side_effect=_http_429_error())) - assert guard._consecutive_rate_limits == 1 + assert guard._consecutive_rate_limits == 2 assert guard._needs_probe is True + assert guard.cooldown_remaining() > 60 @pytest.mark.asyncio diff --git a/ui/lib/api/services/wallet.ts b/ui/lib/api/services/wallet.ts index cbe25dfc..b97e0e72 100644 --- a/ui/lib/api/services/wallet.ts +++ b/ui/lib/api/services/wallet.ts @@ -42,6 +42,7 @@ export interface BalanceDetail { export interface WithdrawResponse { token: string; + mint_url: string; } export interface CreateChildKeyResponse { From 694bc04623892847b3b176be40f4157f13a2384f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 4 Aug 2026 00:06:53 +0200 Subject: [PATCH 46/46] harden melting --- routstr/balance.py | 23 +++++++++ routstr/lightning.py | 22 +++++++-- routstr/payment/lnurl.py | 13 +++++- routstr/wallet.py | 5 +- tests/unit/test_balance.py | 65 ++++++++++++++++++++++++++ tests/unit/test_lnurl_melt_timeout.py | 9 ++-- tests/unit/test_mint_fallback_trust.py | 50 ++++++++++++++++++++ tests/unit/test_periodic_payout.py | 4 +- tests/unit/test_wallet.py | 13 ++++++ 9 files changed, 192 insertions(+), 12 deletions(-) create mode 100644 tests/unit/test_mint_fallback_trust.py diff --git a/routstr/balance.py b/routstr/balance.py index 121e2659..cf050135 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -23,6 +23,7 @@ from .core.db import ( from .core.logging import get_logger from .core.settings import settings from .lightning import lightning_router +from .payment.lnurl import MeltOutcomeAmbiguousError from .wallet import ( classify_redemption_error, credit_balance, @@ -550,6 +551,28 @@ async def refund_wallet_endpoint( }, ) + except MeltOutcomeAmbiguousError as e: + # The melt was dispatched and may still settle. Restoring the balance + # here would let the same debit be paid out twice; keep the debit and + # leave the outcome to reconciliation. + logger.error( + "refund_wallet_endpoint: melt outcome ambiguous; balance withheld " + "pending reconciliation", + extra={ + "error": str(e), + "hashed_key": key.hashed_key, + "remaining_balance": remaining_balance, + "refund_currency": key.refund_currency, + "refund_mint_url": key.refund_mint_url, + }, + ) + raise HTTPException( + status_code=502, + detail=( + "Refund was dispatched but its outcome is unconfirmed; the " + "balance is withheld until reconciliation completes" + ), + ) except HTTPException: # Minting failed — restore the debited balance await _restore_balance( diff --git a/routstr/lightning.py b/routstr/lightning.py index 7a610c10..cb3164bf 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -170,11 +170,23 @@ async def _request_mint_with_fallback( f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}." ) tried: list[str] = [] - candidates = ( - list(dict.fromkeys(allowed_mints)) - if allowed_mints - else _trusted_mint_candidates() - ) + trusted = _trusted_mint_candidates() + if allowed_mints: + # Persisted mint preferences (e.g. an API key's refund_mint_url) must + # not outlive the operator's trusted-mint configuration. + candidates = [m for m in dict.fromkeys(allowed_mints) if m in trusted] + if not candidates: + logger.warning( + "Requested mints are no longer trusted; falling back to " + "configured mints", + extra={ + "requested_mints": list(dict.fromkeys(allowed_mints)), + "op_name": "request_mint_invoice", + }, + ) + candidates = trusted + else: + candidates = trusted for mint_url in candidates: cooldown = mint_cooldown_remaining(mint_url) if cooldown > 0: diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index bb2e6313..f03d1ef5 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -32,6 +32,15 @@ class LNURLError(Exception): """LNURL related errors.""" +class MeltOutcomeAmbiguousError(LNURLError): + """A melt was dispatched but its final outcome could not be confirmed. + + Callers must NOT treat this as a clean failure: the payment may still + settle, so debits backing it must be kept until reconciliation confirms + the true outcome. + """ + + async def decode_lnurl(lnurl: str) -> str: """Decode LNURL to get the actual URL. @@ -268,7 +277,7 @@ async def raw_send_to_lnurl( retry_timeouts=False, ) except Exception as reconciliation_error: - raise LNURLError( + raise MeltOutcomeAmbiguousError( "Melt outcome is ambiguous; quote reconciliation failed and proofs " "must not be retried" ) from reconciliation_error @@ -277,7 +286,7 @@ async def raw_send_to_lnurl( return final_amount state = getattr(getattr(quote, "state", None), "value", "unknown") - raise LNURLError( + raise MeltOutcomeAmbiguousError( "Melt outcome is ambiguous; proofs must not be retried " f"(quote_state={state})" ) from melt_error diff --git a/routstr/wallet.py b/routstr/wallet.py index ef4bf901..486ed8d1 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -1836,7 +1836,10 @@ async def fetch_all_balances( async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: """Send only conservatively proven owner funds for one wallet.""" try: - wallet = await get_wallet(mint_url, unit) + # Runs under wallet_operation_guard; a cached wallet may carry a proof + # snapshot up to 30s stale from another process's reservation, so the + # cross-process lock is only safe with a fresh reload. + wallet = await get_wallet(mint_url, unit, force_reload=True) proofs = get_proofs_per_mint_and_unit( wallet, mint_url, unit, not_reserved=True ) diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index c70fbcb9..1bf94c94 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -781,3 +781,68 @@ async def test_topup_unexpected_non_valueerror_returns_500() -> None: assert exc_info.value.status_code == 500 assert exc_info.value.detail == "Internal server error" + + +@pytest.mark.asyncio +async def test_apikey_refund_ambiguous_melt_does_not_restore_balance() -> None: + """An ambiguous LNURL melt may still settle: the debit must be kept.""" + from fastapi import HTTPException + + from routstr.payment.lnurl import MeltOutcomeAmbiguousError + + key = _make_api_key(balance=5000, refund_address="user@ln.example.com") + + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) + session.commit = AsyncMock() + + with ( + patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), + patch("routstr.balance._refund_cache_set", AsyncMock()), + patch( + "routstr.balance.send_to_lnurl", + AsyncMock(side_effect=MeltOutcomeAmbiguousError("outcome is ambiguous")), + ), + patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore, + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + assert exc_info.value.status_code == 502 + mock_restore.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_apikey_refund_clean_failure_still_restores_balance() -> None: + """A definitively failed melt must keep restoring the debited balance.""" + from fastapi import HTTPException + + key = _make_api_key(balance=5000, refund_address="user@ln.example.com") + + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) + session.commit = AsyncMock() + + with ( + patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), + patch("routstr.balance._refund_cache_set", AsyncMock()), + patch( + "routstr.balance.send_to_lnurl", + AsyncMock(side_effect=RuntimeError("mint rejected melt")), + ), + patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore, + ): + with pytest.raises(HTTPException): + await refund_wallet_endpoint( + authorization="Bearer sk-testhash", + x_cashu=None, + session=session, + ) + + mock_restore.assert_awaited_once() diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index a1b40bcf..b5567570 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -10,7 +10,10 @@ from cashu.core.base import MeltQuoteState from routstr.core.settings import settings from routstr.mint import MintCooldownError, MintRateGuard -from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl +from routstr.payment.lnurl import ( + MeltOutcomeAmbiguousError, + raw_send_to_lnurl, +) LNURL_DATA = { "callback_url": "https://ln.tld/cb", @@ -58,7 +61,7 @@ async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> Non patch.object(settings, "mint_retry_max_attempts", 0), data_patch, invoice_patch, - pytest.raises(LNURLError, match="outcome is ambiguous"), + pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"), ): await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) @@ -106,7 +109,7 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None: patch.object(settings, "mint_operation_timeout_seconds", 5), data_patch, invoice_patch, - pytest.raises(LNURLError, match="outcome is ambiguous"), + pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"), ): await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) diff --git a/tests/unit/test_mint_fallback_trust.py b/tests/unit/test_mint_fallback_trust.py new file mode 100644 index 00000000..f7809458 --- /dev/null +++ b/tests/unit/test_mint_fallback_trust.py @@ -0,0 +1,50 @@ +"""Persisted mint preferences must not bypass the configured trusted set.""" + +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.core.settings import settings +from routstr.lightning import _request_mint_with_fallback + +TRUSTED = "https://good-mint.example.com" +UNTRUSTED = "https://removed-mint.example.com" + + +async def test_untrusted_allowed_mints_fall_back_to_trusted_set() -> None: + attempted: list[str] = [] + + async def fake_get_wallet(mint_url: str, unit: str, **kwargs: object) -> None: + attempted.append(mint_url) + raise ConnectionError("unreachable in test") + + with ( + patch.object(settings, "primary_mint", TRUSTED), + patch.object(settings, "cashu_mints", [TRUSTED]), + patch("routstr.lightning.get_wallet", AsyncMock(side_effect=fake_get_wallet)), + patch("routstr.lightning.mint_cooldown_remaining", return_value=0.0), + ): + with pytest.raises(Exception): + await _request_mint_with_fallback(10, allowed_mints=[UNTRUSTED]) + + assert UNTRUSTED not in attempted + assert attempted == [TRUSTED] + + +async def test_trusted_allowed_mints_are_used_verbatim() -> None: + attempted: list[str] = [] + + async def fake_get_wallet(mint_url: str, unit: str, **kwargs: object) -> None: + attempted.append(mint_url) + raise ConnectionError("unreachable in test") + + with ( + patch.object(settings, "primary_mint", TRUSTED), + patch.object(settings, "cashu_mints", [TRUSTED, "https://other.example.com"]), + patch("routstr.lightning.get_wallet", AsyncMock(side_effect=fake_get_wallet)), + patch("routstr.lightning.mint_cooldown_remaining", return_value=0.0), + ): + with pytest.raises(Exception): + await _request_mint_with_fallback(10, allowed_mints=[TRUSTED]) + + assert attempted == [TRUSTED] diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index 2b4a29fd..c869d7ca 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -147,7 +147,9 @@ async def test_periodic_payout_isolates_failing_mint() -> None: """A failing mint does not prevent payout for the other mints.""" from routstr.core.settings import settings - async def _get_wallet(mint_url: str, unit: str) -> MagicMock: + async def _get_wallet( + mint_url: str, unit: str, force_reload: bool = False + ) -> MagicMock: if mint_url == "http://bad:3338": raise RuntimeError("mint unreachable") return MagicMock() diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index e83e306a..800e4fb9 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -2667,3 +2667,16 @@ async def test_probe_recovery_resets_consecutive_rate_limits() -> None: assert guard._consecutive_rate_limits == 0 assert guard._needs_probe is False assert guard.cooldown_remaining() == 0.0 + + +async def test_payout_reloads_wallet_snapshot_under_guard() -> None: + """Payout must not trust a cached proof snapshot from before the guard.""" + from routstr.wallet import _payout_mint_and_unit + + mock_get_wallet = AsyncMock(side_effect=RuntimeError("stop after get_wallet")) + with patch("routstr.wallet.get_wallet", mock_get_wallet): + await _payout_mint_and_unit("https://mint.example.com", "sat") + + mock_get_wallet.assert_awaited_once_with( + "https://mint.example.com", "sat", force_reload=True + )