diff --git a/.env.example b/.env.example index 9f0c0bbc..6ec85d71 100644 --- a/.env.example +++ b/.env.example @@ -26,9 +26,9 @@ ROUTSTR_SECRET_KEY= # logged at startup. Keep total capacity across all workers below the database # connection limit. Pre-ping is automatic for networked backends; SQLite may # explicitly opt in if desired. -# DATABASE_POOL_SIZE=5 -# DATABASE_MAX_OVERFLOW=10 -# DATABASE_POOL_TIMEOUT=30 +# DATABASE_POOL_SIZE=10 +# DATABASE_MAX_OVERFLOW=20 +# DATABASE_POOL_TIMEOUT=15 # DATABASE_POOL_RECYCLE=1800 # DATABASE_POOL_PRE_PING=false # Warn when a checkout is held this many seconds. @@ -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/api/endpoints.md b/docs/api/endpoints.md index 5a9ea68d..23b019c7 100644 --- a/docs/api/endpoints.md +++ b/docs/api/endpoints.md @@ -327,6 +327,62 @@ GET /v1/models } ``` +### List Model Paths + +Get the selectable upstream routes for each advertised model. This endpoint is +discovery-only; request-side selection will be added separately. + +```http +GET /v1/models/paths +``` + +**Response:** + +```json +{ + "data": [ + { + "id": "anthropic/claude-sonnet-4", + "paths": [ + { + "path": "url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=anthropic%2Fclaude-sonnet-4", + "provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"}, + "endpoint": null + }, + { + "path": "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4&endpoint=google-vertex%2Fus", + "provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"}, + "endpoint": {"tag": "google-vertex/us", "name": "Google"} + } + ] + } + ], + "updated_at": 1753500000 +} +``` + +`path` is an opaque, percent-encoded selector. Clients must store and return it +unchanged rather than parsing or reconstructing it. It identifies the exact +configured route with `url`, `provider-id`, and `model-id`. To avoid exposing +private network details, a configured private IP address or any URL with an +explicit port is advertised as `http://localhost`. OpenRouter routes additionally +preserve the exact machine-readable endpoint `tag`. Provider slugs/types and +endpoint names remain display data. When request-side selection is implemented, +an endpoint tag must not silently fall back to another backend. + +### List Paths for One Model + +Use the exact model ID advertised by `/v1/models`. The query parameter safely +supports IDs containing `/`. + +```http +GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4 +``` + +The response uses the same path objects and `updated_at` field as the collection +endpoint. An unknown model returns `404 Model not found`. A known model whose +paths have not been discovered yet returns `200` with an empty `data` array. + ## Wallet Management ### Create Wallet (Coming Soon) diff --git a/docs/api/overview.md b/docs/api/overview.md index 92fedd7e..c82e4e22 100644 --- a/docs/api/overview.md +++ b/docs/api/overview.md @@ -100,6 +100,7 @@ All errors follow a consistent format: Standard OpenAI-compatible endpoints: - **Models**: `/v1/models` +- **Model paths**: `/v1/models/paths`, `/v1/models/paths/model?model_id=...` - **Responses**: `/v1/responses` - **Chat Completions**: `/v1/chat/completions` - **Embeddings**: `/v1/embeddings` @@ -302,7 +303,7 @@ Get node metadata: GET /v1/info ``` -Supported models and pricing are available at `/v1/models`. +Supported models and pricing are available at `/v1/models`. Upstream provider path discovery is available at `/v1/models/paths` and `/v1/models/paths/model?model_id=...`. ## Next Steps diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index a0358154..68433cc1 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -198,12 +198,23 @@ 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` | | `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` | | `CORS_ORIGINS` | Allowed CORS origins | `*` | | `RELAYS` | Nostr relays (comma-separated) | (default set) | +| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to pause the refresh (previously discovered paths keep being served) | `600` | +| `ENABLE_MODEL_PATHS_REFRESH` | Kill switch for the background model-path refresh (OpenRouter endpoint fan-out) | `true` | + +Mint HTTP 429 responses create a per-mint cooldown. Operations that already hold +Routstr's wallet mutation lock fail fast during that cooldown instead of waiting +while blocking every other wallet mutation. Callers receive an error and may retry +later; the current response does not include the cooldown duration. ### Priority @@ -237,3 +248,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/64ed5594df1f_add_model_paths_table.py b/migrations/versions/64ed5594df1f_add_model_paths_table.py new file mode 100644 index 00000000..a957cc8d --- /dev/null +++ b/migrations/versions/64ed5594df1f_add_model_paths_table.py @@ -0,0 +1,52 @@ +"""add model paths table + +Revision ID: 64ed5594df1f +Revises: aa50fde387a2 +Create Date: 2026-08-02 22:26:33.280409 +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +# revision identifiers, used by Alembic. +revision = "64ed5594df1f" +down_revision = "aa50fde387a2" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "model_paths", + sa.Column("id", sa.Integer(), nullable=False), + sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("provider_slug", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("provider_type", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("endpoint_tag", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("endpoint_name", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("upstream_provider_id", sa.Integer(), nullable=False), + sa.Column("updated_at", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint( + ["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE" + ), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + op.create_index( + op.f("ix_model_paths_upstream_provider_id"), + "model_paths", + ["upstream_provider_id"], + unique=False, + ) + + +def downgrade() -> None: + op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths") + op.drop_table("model_paths") diff --git a/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py b/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py new file mode 100644 index 00000000..7d8abffc --- /dev/null +++ b/migrations/versions/ecfa0d6e2a36_add_mint_url_to_lightning_invoices.py @@ -0,0 +1,67 @@ +"""add mint url to lightning invoices + +Revision ID: ecfa0d6e2a36 +Revises: 64ed5594df1f +Create Date: 2026-08-02 23:53:00.037456 +""" + +import json +import os + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "ecfa0d6e2a36" +down_revision = "64ed5594df1f" +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 74165920..588fc2d1 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -51,6 +51,24 @@ ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900 ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200 +def _format_msat_amount(amount: int) -> str: + sats = f"{amount / 1000:.3f}".rstrip("0").rstrip(".") + return f"{sats} sats ({amount} msats)" + + +def _model_balance_error(required: int, available: int) -> dict[str, dict[str, str]]: + return { + "error": { + "message": ( + f"Insufficient balance: {_format_msat_amount(required)} required " + f"for this model; {_format_msat_amount(available)} available." + ), + "type": "insufficient_quota", + "code": "insufficient_balance", + } + } + + @dataclass(frozen=True) class ReservationSnapshot: release_id: str @@ -264,13 +282,7 @@ async def _validate_bearer_key_locked( ) raise HTTPException( status_code=402, - detail={ - "error": { - "message": f"Insufficient balance: {min_cost} mSats required for this model. {billing_key.total_balance} available.", - "type": "insufficient_quota", - "code": "insufficient_balance", - } - }, + detail=_model_balance_error(min_cost, billing_key.total_balance), ) # Early check: Spending limit check (Child key limit) @@ -360,13 +372,9 @@ async def _validate_bearer_key_locked( if min_cost > 0 and existing_key.total_balance < min_cost: raise HTTPException( status_code=402, - detail={ - "error": { - "message": f"Insufficient balance: {min_cost} mSats required for this model. {existing_key.total_balance} available.", - "type": "insufficient_quota", - "code": "insufficient_balance", - } - }, + detail=_model_balance_error( + min_cost, existing_key.total_balance + ), ) return existing_key @@ -379,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/balance.py b/routstr/balance.py index cf44e612..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, @@ -30,6 +31,7 @@ from .wallet import ( recieve_token, send_to_lnurl, send_token, + token_mint_url, ) router = APIRouter() @@ -184,6 +186,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, @@ -201,6 +214,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.info( + "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: @@ -209,12 +234,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.info( + "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} @@ -290,7 +344,11 @@ async def _get_persisted_api_key_refund( 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 = ( @@ -305,7 +363,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, + }, ) @@ -450,15 +512,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, @@ -468,10 +529,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": @@ -490,13 +551,47 @@ 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(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", @@ -523,7 +618,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", @@ -717,7 +812,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 38112efb..604ed948 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, @@ -51,6 +45,13 @@ ADMIN_SESSION_DURATION = 3600 MAX_USAGE_ANALYTICS_HOURS = 365 * 24 +async def _refresh_provider_model_paths(upstream_provider_id: int) -> None: + """Queue discovery sync without blocking the committed admin mutation.""" + from ..upstream.model_paths import schedule_model_paths_refresh_for_provider + + await schedule_model_paths_refresh_for_provider(upstream_provider_id) + + async def require_admin_api(request: Request) -> None: auth_header = request.headers.get("Authorization") if not auth_header or not auth_header.startswith("Bearer "): @@ -435,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", @@ -476,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): @@ -579,6 +574,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 +629,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 +649,7 @@ async def delete_all_provider_models(provider_id: str) -> dict[str, object]: await session.delete(row) # type: ignore await session.commit() await refresh_model_maps() + await _refresh_provider_model_paths(provider_pk) return {"ok": True, "deleted": len(rows)} @@ -743,6 +741,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, @@ -1032,6 +1031,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) @@ -1057,6 +1057,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) @@ -1092,6 +1093,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 1fc70400..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 = ( @@ -317,9 +319,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in .where(col(ApiKey.total_spent) == 0) .where(col(ApiKey.total_requests) == 0) .where(col(ApiKey.parent_key_hash).is_(None)) - .where( - (col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff) - ) + .where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)) .where(~pending_invoice) .where(~has_children) ) @@ -373,6 +373,60 @@ class ModelRow(SQLModel, table=True): # type: ignore upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models") +class ModelPathRow(SQLModel, table=True): # type: ignore + """Upstream provider path a model is reachable through. + + Discovery/visibility data only. ``model_id`` is intentionally NOT globally + unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or + id``) grouped across every provider that exposes the model. A single model + can therefore have several rows — one per direct provider path plus one per + OpenRouter sub-provider endpoint. + """ + + __tablename__ = "model_paths" + __table_args__ = ( + UniqueConstraint( + "model_id", + "path", + "upstream_provider_id", + name="uq_model_paths_model_path_provider", + ), + ) + id: int | None = Field(default=None, primary_key=True) + # No standalone index on model_id: the unique constraint's autoindex already + # leads on model_id, so a second index only adds write amplification. + model_id: str = Field( + description="Client-visible /v1/models id (forwarded_model_id or id)" + ) + path: str = Field( + description=( + "Opaque selector containing upstream URL, provider ID, model ID, " + "and optional endpoint tag" + ) + ) + provider_slug: str = Field( + description="Public slug of the configured upstream provider" + ) + provider_type: str = Field(description="Configured upstream provider type") + endpoint_tag: str | None = Field( + default=None, + description="Exact OpenRouter endpoint tag used for request-side selection", + ) + endpoint_name: str | None = Field( + default=None, description="Human-readable endpoint display name" + ) + upstream_provider_id: int = Field( + index=True, + foreign_key="upstream_providers.id", + ondelete="CASCADE", + description="upstream_providers.id this path was discovered from", + ) + updated_at: int = Field( + default=0, + description="Unix timestamp of the refresh cycle that wrote this row", + ) + + class LightningInvoice(SQLModel, table=True): # type: ignore __tablename__ = "lightning_invoices" @@ -383,12 +437,18 @@ 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" ) 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" ) @@ -663,9 +723,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())) @@ -764,9 +822,7 @@ async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool: return result.rowcount == 1 -async def complete_routstr_fee_payout( - session: AsyncSession, paid_msats: int -) -> bool: +async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) -> bool: """Mark a checkpointed payout complete after the external payment succeeds.""" stmt = ( update(RoutstrFee) diff --git a/routstr/core/main.py b/routstr/core/main.py index ca1a5a93..979f5cdb 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -58,6 +58,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: providers_task = None models_refresh_task = None model_maps_refresh_task = None + model_paths_refresh_task = None key_reset_task = None stale_reservation_task = None dead_key_prune_task = None @@ -130,6 +131,13 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: refresh_upstreams_models_periodically(get_upstreams) ) model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) + # Always started: the loop re-reads the enable flag and interval every + # iteration, so 0 -> N (or re-enabling) takes effect without a restart. + from ..upstream.model_paths import refresh_model_paths_periodically + + model_paths_refresh_task = asyncio.create_task( + refresh_model_paths_periodically(get_upstreams) + ) payout_task = asyncio.create_task(periodic_payout()) if global_settings.nsec: nip91_task = asyncio.create_task(announce_provider()) @@ -137,9 +145,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: if global_settings.providers_refresh_interval_seconds > 0: providers_task = asyncio.create_task(providers_cache_refresher()) key_reset_task = asyncio.create_task(periodic_key_reset()) - stale_reservation_task = asyncio.create_task( - periodic_stale_reservation_sweep() - ) + stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep()) dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) @@ -176,6 +182,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: models_refresh_task.cancel() if model_maps_refresh_task is not None: model_maps_refresh_task.cancel() + if model_paths_refresh_task is not None: + model_paths_refresh_task.cancel() if key_reset_task is not None: key_reset_task.cancel() if stale_reservation_task is not None: @@ -209,6 +217,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(models_refresh_task) if model_maps_refresh_task is not None: tasks_to_wait.append(model_maps_refresh_task) + if model_paths_refresh_task is not None: + tasks_to_wait.append(model_paths_refresh_task) if key_reset_task is not None: tasks_to_wait.append(key_reset_task) if stale_reservation_task is not None: @@ -245,9 +255,7 @@ class _ImmutableStaticFiles(StaticFiles): async def get_response(self, path: str, scope: Scope) -> StarletteResponse: response = await super().get_response(path, scope) if response.status_code == 200: - response.headers["Cache-Control"] = ( - "public, max-age=31536000, immutable" - ) + response.headers["Cache-Control"] = "public, max-age=31536000, immutable" return response @@ -321,9 +329,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir(): # Serve the App Router RSC payload for the home page. @app.get("/index.txt", include_in_schema=False) async def serve_root_rsc() -> FileResponse: - return FileResponse( - UI_DIST_PATH / "index.txt", media_type="text/x-component" - ) + return FileResponse(UI_DIST_PATH / "index.txt", media_type="text/x-component") # Next.js is built with `trailingSlash: true`, so all UI page URLs end # with a slash (e.g. `/login/`). The proxy router catches `/{path:path}` diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 0caffd16..6852eaa2 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -53,6 +53,18 @@ 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 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") # Pricing # Default behavior: derive pricing from MODELS @@ -98,22 +110,30 @@ class Settings(BaseSettings): models_refresh_interval_seconds: int = Field( default=360, env="MODELS_REFRESH_INTERVAL_SECONDS" ) + model_paths_refresh_interval_seconds: int = Field( + default=600, env="MODEL_PATHS_REFRESH_INTERVAL_SECONDS" + ) enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") + enable_model_paths_refresh: bool = Field( + default=True, env="ENABLE_MODEL_PATHS_REFRESH" + ) refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") - refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS") + refund_sweep_ttl_seconds: int = Field( + default=604800, env="REFUND_SWEEP_TTL_SECONDS" + ) refund_sweep_claim_timeout_seconds: int = Field( 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") @@ -138,9 +158,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.""" diff --git a/routstr/lightning.py b/routstr/lightning.py index d75a8c2c..cb3164bf 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -1,8 +1,13 @@ import asyncio import hashlib +import re import secrets import time +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 @@ -12,12 +17,85 @@ 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, wallet_operation_guard +from .mint import ( + is_mint_rate_limited, + mint_cooldown_remaining, + run_mint_operation, +) +from .wallet import ( + MintConnectionError, + get_wallet, + is_mint_connection_error, + wallet_operation_guard, +) logger = get_logger(__name__) lightning_router = APIRouter(prefix="/lightning") +# 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 + 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) +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, + ) + + +def _publish_invoice_value(invoice: LightningInvoice, key: str, value: Any) -> None: + """Update a caller view without marking a mapped object dirty.""" + try: + set_committed_value(invoice, key, value) + except AttributeError: + setattr(invoice, key, value) + class InvoiceCreateRequest(BaseModel): amount_sats: int = Field(gt=0, le=1_000_000, description="Amount in satoshis") @@ -61,16 +139,102 @@ class InvoiceStatusResponse(BaseModel): expires_at: int +_RETRYABLE_INVOICE_STATUSES = ("pending", "settlement_pending") + + 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, + *, + allowed_mints: list[str] | None = None, +) -> tuple[str, str, str]: + """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] = [] + 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: + tried.append(f"{mint_url}: cooling down") + logger.info( + "Skipping mint during cooldown", + extra={ + "mint_url": mint_url, + "cooldown_seconds": round(cooldown, 2), + "op_name": "request_mint_invoice", + }, + ) + continue + try: + wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False) + quote = await run_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: + 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 + 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, allowed_mints=allowed_mints + ) + return bolt11, payment_hash, mint_url def generate_invoice_id() -> str: @@ -84,6 +248,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: @@ -94,14 +259,23 @@ 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" - bolt11, payment_hash = await generate_lightning_invoice( - request.amount_sats, description + allowed_mints = None + if request.purpose == "topup": + assert topup_api_key is not None + # A key's liabilities are attributed to a single refund mint. Keep + # top-up collateral on that same mint so balances and payouts cannot + # misclassify funds held by another mint as owner profit. + 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, allowed_mints=allowed_mints ) invoice_id = generate_invoice_id() @@ -116,6 +290,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, @@ -161,12 +336,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": @@ -200,8 +375,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": @@ -220,161 +399,256 @@ 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: - # Minting makes proofs visible before database finalization. Share the - # cross-process wallet guard with owner payout so that visibility and the - # corresponding liability commit are observed atomically by the payout loop. - async with wallet_operation_guard(): - await _check_invoice_payment_locked(invoice, session) +) -> 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 not in _RETRYABLE_INVOICE_STATUSES: + await session.commit() + return False + observed_status = invoice.status + settlement = _InvoiceSettlement.from_invoice(invoice) + await session.commit() -async def _check_invoice_payment_locked( - invoice: LightningInvoice, session: AsyncSession -) -> None: - minted = False - invoice_id = invoice.id - invoice_purpose = invoice.purpose - invoice_amount_sats = invoice.amount_sats - invoice_payment_hash = invoice.payment_hash - finalized_api_key_hash = invoice.api_key_hash - try: - # A preceding invoice lookup starts a transaction. End it before the - # potentially slow mint request so it cannot pin a pool connection. - await session.commit() + mint_url = settlement.mint_url or settings.primary_mint + wallet = await get_wallet(mint_url, "sat") + mint_status = await run_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 getattr(mint_status, "state", None) == MintQuoteState.unpaid + payment_confirmed = True - wallet = await get_wallet(settings.primary_mint, "sat") - mint_status = await wallet.get_mint_quote(invoice_payment_hash) - if not mint_status.paid: - return + # 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 - # Do not redeem a paid top-up quote if its target has already been - # pruned. This validation owns a short-lived session and releases its - # connection before mint redemption starts. - if invoice_purpose == "topup": - if not finalized_api_key_hash: - raise ValueError("No API key associated with topup invoice") - async with create_session() as validation_session: - target = await validation_session.get(ApiKey, finalized_api_key_hash) - if target is None: - terminal = await validation_session.exec( # type: ignore[call-overload] - update(LightningInvoice) - .where( - col(LightningInvoice.id) == invoice_id, - col(LightningInvoice.status) == "pending", - ) - .values(status="reconciliation_required") + # Reject a paid top-up whose target was pruned before redeeming its + # single-use quote. The validation session is closed before mint I/O. + if settlement.purpose == "topup": + if not settlement.api_key_hash: + raise ValueError("No API key associated with topup invoice") + async with create_session() as validation_session: + target = await validation_session.get( + ApiKey, settlement.api_key_hash ) - await validation_session.commit() - if terminal.rowcount == 1: - set_committed_value( - invoice, "status", "reconciliation_required" - ) - else: - committed_invoice = await validation_session.get( - LightningInvoice, invoice_id - ) - if committed_invoice is not None: - set_committed_value( - invoice, "status", committed_invoice.status + if target is None: + terminal = await validation_session.exec( # type: ignore[call-overload] + update(LightningInvoice) + .where( + col(LightningInvoice.id) == settlement.id, + col(LightningInvoice.status).in_( + _RETRYABLE_INVOICE_STATUSES + ), ) - logger.critical( - "Paid topup invoice target API key was not found; reconciliation required", - extra={"invoice_id": invoice_id}, - ) - return + .values(status="reconciliation_required") + ) + await validation_session.commit() + if terminal.rowcount == 1: + _publish_invoice_value( + invoice, "status", "reconciliation_required" + ) + else: + await _reload_invoice_view(invoice, session) + logger.critical( + "Paid topup invoice target API key was not found; reconciliation required", + extra={"invoice_id": settlement.id}, + ) + return False - # The mint enforces single-use quotes, so a concurrent checker that - # races us here fails inside wallet.mint rather than double-minting. - await wallet.mint(invoice_amount_sats, quote_id=invoice_payment_hash) - minted = True + # Quote-linked proof verification makes an ambiguous mint response + # retryable without crediting unrelated wallet balance growth. + await _mint_invoice_quote(wallet, settlement) + minted = True - # Paid finalization owns a fresh session. The API/watcher session is - # never rolled back by this function, so its invoice and sibling ORM - # objects remain usable after a DB failure or lost CAS race. - async with create_session() as finalization_session: - if invoice_purpose == "create": - api_key = await _create_api_key_record(invoice, finalization_session) - finalized_api_key_hash = api_key.hashed_key - elif invoice_purpose == "topup": - await _credit_topup_record(invoice, finalization_session) - - # Conditional transition guards against double-credit: the credit - # above and this status flip commit atomically, and a lost race - # rolls both back in the owned finalization session. paid_at = int(time.time()) - finalized = await finalization_session.exec( # type: ignore[call-overload] - update(LightningInvoice) - .where( - col(LightningInvoice.id) == invoice_id, - col(LightningInvoice.status) == "pending", - ) - .values( - status="paid", - paid_at=paid_at, - api_key_hash=finalized_api_key_hash, + async with create_session() as finalization_session: + settled, api_key_hash = await _finalize_invoice_settlement( + settlement, finalization_session, paid_at ) + if not settled: + await _reload_invoice_view(invoice, session) + return False + + _publish_invoice_value(invoice, "status", "paid") + _publish_invoice_value(invoice, "paid_at", paid_at) + _publish_invoice_value(invoice, "api_key_hash", api_key_hash) + logger.info( + "Lightning invoice paid", + extra={ + "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, + }, ) - if finalized.rowcount != 1: - await finalization_session.rollback() - committed_invoice = await finalization_session.get( - LightningInvoice, invoice_id - ) - await finalization_session.commit() - if committed_invoice is not None: - # A concurrent finalizer won the CAS. Publish only the - # state observed from the database after ending the owned - # read transaction; never refresh the caller's session. - set_committed_value( - invoice, "api_key_hash", committed_invoice.api_key_hash + 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)}, ) - set_committed_value(invoice, "status", committed_invoice.status) - set_committed_value(invoice, "paid_at", committed_invoice.paid_at) - return - await finalization_session.commit() + if minted: + logger.critical( + "Invoice mint succeeded but DB finalization failed; reconciliation required", + extra={"invoice_id": invoice.id, "purpose": invoice.purpose}, + ) + try: + await _reload_invoice_view(invoice, session) + except Exception: + pass + if not isinstance(error, Exception): + raise + logger.error(f"Failed to check invoice payment: {error}") + return False - # Only publish finalized values to the caller-owned object after the - # owned transaction has committed successfully. - set_committed_value(invoice, "api_key_hash", finalized_api_key_hash) - set_committed_value(invoice, "status", "paid") - set_committed_value(invoice, "paid_at", paid_at) - logger.info( - "Lightning invoice paid", - extra={ - "invoice_id": invoice_id, - "amount_sats": invoice_amount_sats, - "purpose": invoice_purpose, - "api_key_hash": finalized_api_key_hash[:8] + "..." - if finalized_api_key_hash - else None, - }, +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, ) - except BaseException as e: - # BaseException so task cancellation (e.g. client disconnect) after a - # successful mint still triggers the reconciliation alert. Any rollback - # belongs to create_session(), never to the caller-owned session. - if minted: - logger.critical( - "Invoice mint succeeded but DB finalization failed; reconciliation required", - extra={"invoice_id": invoice_id, "purpose": invoice_purpose}, - ) - if not isinstance(e, Exception): + 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 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, + retry_timeouts=False, + ) + except Exception as error: + if not _is_outputs_already_signed(error): raise - logger.error(f"Failed to check invoice payment: {e}") + + 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_amount = _invoice_quote_proof_amount(wallet, invoice.payment_hash) + if minted_amount < invoice.amount_sats: + raise RuntimeError( + "Invoice mint succeeded but quote-linked proofs total " + f"{minted_amount} 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, session: AsyncSession + invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession ) -> ApiKey: - dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" - hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest() + mint_url = invoice.mint_url or settings.primary_mint api_key = ApiKey( - hashed_key=hashed_key, + hashed_key=_invoice_api_key_hash(invoice), balance=invoice.amount_sats * 1000, 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, @@ -384,50 +658,142 @@ async def _create_api_key_record( return api_key -async def _credit_topup_record( - invoice: LightningInvoice, session: AsyncSession +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") - credited = await session.exec( # type: ignore[call-overload] + 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 credited.rowcount != 1: + if result.rowcount != 1: raise ValueError("Associated API key not found") -INVOICE_WATCH_INTERVAL_SECONDS = 5 +async def _finalize_invoice_settlement( + invoice: _InvoiceSettlement, session: AsyncSession, paid_at: int +) -> tuple[bool, str | None]: + """Atomically fence and apply one invoice credit in the provided owned session.""" + 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).in_(_RETRYABLE_INVOICE_STATUSES) + ) + .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 + + 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() + return True, api_key_hash + + +async def _reload_invoice_view( + invoice: LightningInvoice, _caller_session: AsyncSession +) -> None: + """Publish committed invoice state without touching the caller transaction.""" + async with create_session() as reload_session: + stored = await reload_session.get(LightningInvoice, invoice.id) + if stored is None: + return + status = stored.status + paid_at = stored.paid_at + api_key_hash = stored.api_key_hash + await reload_session.commit() + _publish_invoice_value(invoice, "status", status) + _publish_invoice_value(invoice, "paid_at", paid_at) + _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: + await _topup_api_key_record(invoice, session) + + +# 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 -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 new file mode 100644 index 00000000..c0f0ea3c --- /dev/null +++ b/routstr/mint.py @@ -0,0 +1,343 @@ +"""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: + 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 + + 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: + 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: + 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) + self.apply_rate_limit_cooldown(retry_after) + 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: + 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: + 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/cost_calculation.py b/routstr/payment/cost_calculation.py index 37ac15d3..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,6 +509,38 @@ 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 + ) + + # 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 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( + round( + raw_input_msats * cache_creation_weight / total_input_weight, + 3, + ) + ) logger.info( "Using cost from usage data/details", @@ -487,6 +548,8 @@ def _calculate_from_usd_cost( "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 +564,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/payment/helpers.py b/routstr/payment/helpers.py index afd71219..702ab527 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -18,7 +18,6 @@ from ..wallet import deserialize_token_from_string logger = get_logger(__name__) - 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 @@ -243,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 3f07e36c..f03d1ef5 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -1,16 +1,17 @@ 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. Bound it here so callers fail instead of hanging forever. -MELT_TIMEOUT_SECONDS = 60 +from ..mint import ( + MINT_TRANSPORT_EXCEPTIONS, + is_mint_rate_limited, + run_mint_operation, +) try: from bech32 import bech32_decode, convertbits # type: ignore @@ -31,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. @@ -221,23 +231,62 @@ 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 run_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) try: - _ = await asyncio.wait_for( - wallet.melt( + 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 asyncio.TimeoutError as e: - raise LNURLError( - f"Melt timed out after {MELT_TIMEOUT_SECONDS}s (mint unresponsive)" - ) from e - return final_amount + 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: + 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 MeltOutcomeAmbiguousError( + "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 MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous; proofs must not be retried " + f"(quote_state={state})" + ) from melt_error diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 17e8ac64..5c634ced 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -455,7 +455,9 @@ async def _update_sats_pricing_once() -> None: for m in upstream.get_cached_models() ] upstream._models_cache = updated_models - upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models} + upstream._models_by_id = { + m.forwarded_model_id or m.id: m for m in updated_models + } updated_count += len(updated_models) if updated_count > 0: @@ -510,9 +512,7 @@ class ModelTestRequest(V2BaseModel): request_data: dict -@models_router.post( - "/api/models/test", dependencies=[Depends(_require_admin_api)] -) +@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)]) async def test_model( payload: ModelTestRequest, session: AsyncSession = Depends(get_session), @@ -595,6 +595,37 @@ async def test_model( } +@models_router.get("/v1/models/paths") +@models_router.get("/v1/models/paths/", include_in_schema=False) +async def model_paths() -> dict: + """All models with every upstream provider path they are reachable through.""" + from ..upstream.model_paths import get_all_model_paths + + return await get_all_model_paths() + + +@models_router.get("/v1/models/paths/model") +@models_router.get("/v1/models/paths/model/", include_in_schema=False) +async def model_paths_for_model(model_id: str) -> dict: + """Paths for a single model. + + Uses a query parameter (``?model_id=...``) under a fully static route so + model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL + encoding and there is no dynamic-route ambiguity. + """ + from ..proxy import get_unique_models + from ..upstream.model_paths import get_paths_for_model + + result = await get_paths_for_model(model_id) + if not result["data"]: + advertised_ids = { + model.forwarded_model_id or model.id for model in get_unique_models() + } + if model_id not in advertised_ids: + raise HTTPException(status_code=404, detail="Model not found") + return result + + @models_router.get("/v1/models") @models_router.get("/v1/models/", include_in_schema=False) @models_router.get("/models") diff --git a/routstr/proxy.py b/routstr/proxy.py index 887eb38c..4542d008 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -26,7 +26,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 ( calculate_discounted_max_cost, check_token_balance, @@ -191,6 +190,19 @@ async def refresh_model_maps() -> None: disabled_model_keys=disabled_model_keys, ) + # Keep model-path discovery in sync with admin mutations: disabling or + # deleting a provider must stop advertising its paths immediately rather + # than after the next timed refresh. + from .upstream.model_paths import prune_model_paths_for_inactive_providers + + try: + await prune_model_paths_for_inactive_providers() + except Exception as e: # noqa: BLE001 - discovery sync must not break routing + logger.warning( + "Failed to prune model paths for inactive providers", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + async def refresh_model_maps_periodically() -> None: """Background task to refresh model maps every minute.""" @@ -354,8 +366,6 @@ async def _proxy( max_cost_for_model = await calculate_discounted_max_cost( _max_cost_for_model, request_body_dict, model_obj=model_obj ) - # Ensure max_cost_for_model is at least the minimum allowed request cost - max_cost_for_model = max(max_cost_for_model, settings.min_request_msat) check_token_balance(headers, request_body_dict, max_cost_for_model) @@ -494,7 +504,6 @@ 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) if candidate_max > max_cost_for_model: await revert_pay_for_request( key, session, max_cost_for_model, reservation_snapshot @@ -794,17 +803,29 @@ async def get_bearer_token_key( }, ) return key - except Exception as e: - key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key - logger.error( - f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}", + except HTTPException as error: + detail: dict[str, Any] = error.detail if isinstance(error.detail, dict) else {} + raw_error = detail.get("error") + error_info = raw_error if isinstance(raw_error, dict) else {} + logger.warning( + "Bearer token rejected", extra={ - "error": str(e), - "error_type": type(e).__name__, + "status_code": error.status_code, + "error_code": error_info.get("code"), "path": path, "model_id": model_id, - "min_cost_msat": min_cost, - "bearer_key_preview": key_preview, + "required_msat": min_cost, + }, + ) + raise + except Exception as error: + logger.exception( + "Bearer token validation failed", + extra={ + "error_type": type(error).__name__, + "path": path, + "model_id": model_id, + "required_msat": min_cost, }, ) raise diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index 78b7066c..f6e46d59 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -26,7 +26,9 @@ from ..wallet import ( execute_bolt11_payment, maximum_owner_cashu_balance_sats, prepare_bolt11_payment, + release_token_reservation, send_token, + token_mint_url, ) from .ppqai import PPQAIUpstreamProvider from .routstr import RoutstrUpstreamProvider @@ -264,12 +266,13 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) return + actual_mint_url = token_mint_url(token, mint_url) try: 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", @@ -277,8 +280,24 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: except Exception: logger.critical( "Aborting auto top-up because its cashu token could not be persisted", - extra={"provider_id": row.id, "mint_url": mint_url}, + extra={"provider_id": row.id, "mint_url": actual_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 a8dba7e3..71f337ca 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -53,6 +53,7 @@ from ..wallet import ( classify_redemption_error, recieve_token, send_token, + token_mint_url, ) from . import messages_dispatch from .cache_breakpoints import ( @@ -69,6 +70,81 @@ 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: CostMetadata +) -> 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( + 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: CostMetadata) -> 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": 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") + ), + } + 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"] = int(_cost_field(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: @@ -275,30 +351,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 - # Inject into top-level usage block (OpenAI/Anthropic style) - if "usage" in response_json: - response_json["usage"]["cost"] = total_usd - 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", {}) @@ -1229,12 +1299,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 ) @@ -1281,6 +1348,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 @@ -1666,12 +1734,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 ) @@ -1718,6 +1783,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 @@ -2153,6 +2219,9 @@ class BaseUpstreamProvider: if k.lower() in allowed_headers } + # 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(), status_code=response.status_code, @@ -2242,9 +2311,14 @@ class BaseUpstreamProvider: ) self.inject_cost_metadata(response_json, cost_data, key) + # Inject the same cost headers used by every paid response path. + response_headers: dict[str, str] = {} + _inject_cost_response_headers(response_headers, cost_data) + return Response( content=json.dumps(response_json).encode(), status_code=200, + headers=response_headers, media_type="application/json", ) @@ -2295,11 +2369,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 ) @@ -2545,7 +2620,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 @@ -2573,7 +2648,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", @@ -2600,6 +2675,7 @@ class BaseUpstreamProvider: }, ) + cost_data: CostData | MaxCostData | None = None if ( input_tokens > 0 or output_tokens > 0 @@ -2655,9 +2731,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(), @@ -3520,7 +3617,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, ) @@ -3700,6 +3797,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 +3824,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 +3883,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 +3917,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": @@ -3873,7 +3984,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, ) @@ -4681,6 +4792,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 +4819,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 +4867,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 +4898,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": @@ -4843,7 +4965,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/upstream/model_paths.py b/routstr/upstream/model_paths.py new file mode 100644 index 00000000..95b6edf2 --- /dev/null +++ b/routstr/upstream/model_paths.py @@ -0,0 +1,819 @@ +"""Model-path discovery service. + +Exposes every selectable upstream route a Routstr model is reachable through. +This PR remains discovery-only: request-side routing will consume the opaque +selectors in a follow-up. + +A path is a standard percent-encoded query string containing the configured +upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter +endpoint, its machine-readable tag:: + + url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=claude-sonnet-4 + url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=claude-sonnet-4&endpoint=google-vertex%2Fus +""" + +from __future__ import annotations + +import asyncio +import ipaddress +import random +import time +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable +from urllib.parse import urlencode, urlsplit + +import httpx +from sqlalchemy.dialects.sqlite import insert +from sqlalchemy.orm import selectinload +from sqlmodel import col, delete, select + +from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session +from ..core.logging import get_logger + +if TYPE_CHECKING: + from sqlmodel.ext.asyncio.session import AsyncSession + + from .base import BaseUpstreamProvider + +logger = get_logger(__name__) + +# Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds +# of models does not open hundreds of concurrent requests every refresh. +_OPENROUTER_CONCURRENCY = 5 +_OPENROUTER_TIMEOUT_SECONDS = 10.0 + +# Rows inserted per statement during persist. Keeps each INSERT bounded while +# avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle. +_PERSIST_CHUNK_SIZE = 500 + +# Admin mutations enqueue provider IDs here instead of running OpenRouter's +# per-model endpoint fan-out inside the request. One worker serializes refreshes +# and coalesces repeated mutations for the same provider. +_scheduled_provider_refresh_ids: set[int] = set() +_scheduled_provider_refresh_task: asyncio.Task[None] | None = None + +# Visibility key used across this module: routing carries the provider +# dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)), +# so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps. +ModelKey = tuple[str, int] + + +@dataclass(frozen=True) +class EndpointIdentity: + """Exact OpenRouter endpoint identity returned by ``/endpoints``.""" + + tag: str + provider_name: str | None + + +@dataclass(frozen=True) +class ConfiguredProviderIdentity: + """Public identity of one configured upstream provider.""" + + id: int + slug: str + provider_type: str + base_url: str + + +@dataclass(frozen=True) +class DiscoveredPath: + """One model route ready for persistence and API serialization.""" + + model_id: str + path: str + provider: ConfiguredProviderIdentity + endpoint_tag: str | None = None + endpoint_name: str | None = None + + +@dataclass(frozen=True) +class ProviderPathSnapshot: + """Refresh result plus model IDs whose prior rows must survive degradation.""" + + paths: tuple[DiscoveredPath, ...] + preserve_model_ids: frozenset[str] = frozenset() + + +def public_provider_url(base_url: str) -> str: + """Mask private IP addresses and URLs with explicit ports.""" + parsed = urlsplit(base_url) + try: + if parsed.port is not None: + return "http://localhost" + except ValueError: + # An invalid explicit port must not accidentally leak through. + return "http://localhost" + + hostname = parsed.hostname + if hostname is None: + return base_url + try: + address = ipaddress.ip_address(hostname) + except ValueError: + return base_url + return "http://localhost" if address.is_private else base_url + + +def encode_model_path( + base_url: str, + provider_id: int, + model_id: str, + endpoint_tag: str | None = None, +) -> str: + """Encode the complete upstream route selector advertised to clients.""" + components: list[tuple[str, str | int]] = [ + ("url", base_url), + ("provider-id", provider_id), + ("model-id", model_id), + ] + if endpoint_tag: + components.append(("endpoint", endpoint_tag)) + return urlencode(components) + + +def _make_http_client() -> httpx.AsyncClient: + """Client factory, separated so tests can substitute a mock transport.""" + return httpx.AsyncClient() + + +def is_openrouter_base_url(base_url: str | None) -> bool: + """True when ``base_url`` points at OpenRouter. + + Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``: + that predicate also returns True for native Anthropic (correct for + cache-control, wrong for OpenRouter endpoint discovery). This one keys only + on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched + while native Anthropic is not. + """ + return "openrouter.ai" in (base_url or "") + + +def exposed_model_id(model: object) -> str: + """Return exactly the ID advertised by ``/v1/models``. + + A forwarded ID is already a public routable alias and must remain intact, + including any slash. Without one, ``/v1/models`` exposes the base ID. + """ + forwarded = getattr(model, "forwarded_model_id", None) + if forwarded: + return forwarded + return public_model_id(getattr(model, "id")) + + +def public_model_id(model_id: str) -> str: + """Model id exposed by model-path API responses. + + Uses the same rule as ``create_model_mappings.get_base_model_id`` and + ``resolve_model_alias`` — strip everything before the *first* slash — so + the id shown here can be sent back to ``/v1/chat/completions`` verbatim. + """ + return model_id.split("/", 1)[1] if "/" in model_id else model_id + + +def openrouter_author_slug(model: object) -> str | None: + """Return a canonical ``author/slug`` for the OpenRouter endpoints API. + + Prefer ``canonical_slug``, then a slash-containing ``id``, then a + slash-containing ``forwarded_model_id``. The forwarded id is exactly what + the proxy sends upstream for admin-created alias rows (``base.py`` forwards + ``forwarded_model_id or id``), so it is a valid OpenRouter id when the + bare ``id`` is a local alias with no slash. + """ + canonical = getattr(model, "canonical_slug", None) + if canonical and "/" in canonical: + return canonical + model_id = getattr(model, "id", None) + if model_id and "/" in model_id: + return model_id + forwarded = getattr(model, "forwarded_model_id", None) + if forwarded and "/" in forwarded: + return forwarded + return None + + +class _RefreshCycleState: + """Per-refresh shared state: fetch dedupe cache and rate-limit latch. + + ``endpoint_cache`` dedupes byte-identical ``/endpoints`` fetches when two + providers point at the same OpenRouter base URL. ``rate_limited`` latches + on the first 429 so the rest of the cycle stops hammering a throttled API; + the whole provider result then degrades to "unknown" instead of an empty + list, which preserves previously persisted rows. + """ + + def __init__(self) -> None: + self.endpoint_cache: dict[tuple[str, str], list[EndpointIdentity] | None] = {} + self.rate_limited = False + + +async def _fetch_openrouter_endpoint_subproviders( + client: httpx.AsyncClient, + base_url: str, + api_key: str, + author_slug: str, + semaphore: asyncio.Semaphore, + cycle: _RefreshCycleState, +) -> list[EndpointIdentity] | None: + """Return exact endpoint identities for one model, or ``None`` when unknown. + + ``None`` (not ``[]``) signals a degraded fetch — network failure, rate + limit, non-200, or an unparseable payload — so callers can distinguish + "this model has no endpoints" from "we could not find out". Failures are + logged and swallowed so one model never breaks the whole refresh. + """ + cache_key = (base_url, author_slug) + if cache_key in cycle.endpoint_cache: + return cycle.endpoint_cache[cache_key] + if cycle.rate_limited: + return None + + url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints" + headers = {"Authorization": f"Bearer {api_key}"} if api_key else {} + result: list[EndpointIdentity] | None + async with semaphore: + try: + resp = await client.get( + url, headers=headers, timeout=_OPENROUTER_TIMEOUT_SECONDS + ) + except Exception as e: # noqa: BLE001 - isolate per-model failures + logger.warning( + "OpenRouter endpoint discovery request failed", + extra={"author_slug": author_slug, "error": str(e)}, + ) + cycle.endpoint_cache[cache_key] = None + return None + + if resp.status_code == 429: + logger.warning( + "OpenRouter endpoint discovery rate-limited; aborting cycle", + extra={"author_slug": author_slug}, + ) + cycle.rate_limited = True + cycle.endpoint_cache[cache_key] = None + return None + if resp.status_code != 200: + logger.warning( + "OpenRouter endpoint discovery non-200", + extra={"author_slug": author_slug, "status_code": resp.status_code}, + ) + cycle.endpoint_cache[cache_key] = None + return None + + try: + payload = resp.json() + data = payload.get("data") if isinstance(payload, dict) else None + endpoints = data.get("endpoints") if isinstance(data, dict) else None + if not isinstance(endpoints, list): + raise ValueError("endpoints must be a list") + identities: dict[str, EndpointIdentity] = {} + for endpoint in endpoints: + if not isinstance(endpoint, dict): + continue + tag = endpoint.get("tag") + if not isinstance(tag, str) or not tag.strip(): + continue + provider_name = endpoint.get("provider_name") + identities.setdefault( + tag, + EndpointIdentity( + tag=tag, + provider_name=provider_name + if isinstance(provider_name, str) and provider_name + else None, + ), + ) + if endpoints and not identities: + raise ValueError("endpoints contain no usable tags") + result = list(identities.values()) + except Exception as e: # noqa: BLE001 + logger.warning( + "OpenRouter endpoint discovery bad payload", + extra={"author_slug": author_slug, "error": str(e)}, + ) + result = None + + cycle.endpoint_cache[cache_key] = result + return result + + +async def _load_model_visibility() -> tuple[ + dict[ModelKey, ModelRow], + set[ModelKey], + dict[int, ConfiguredProviderIdentity], +]: + """Load the same DB model visibility inputs used by routing. + + ``refresh_model_maps`` builds routing from enabled providers, enabled DB + override rows, and disabled model keys — all keyed on + ``(model_id.lower(), upstream_provider_id)`` because ``ModelRow``'s primary + key is composite and the same id legitimately exists on several providers. + Model-path discovery uses the same keying so disabling a model on one + provider never hides it on another, and one provider's + ``forwarded_model_id`` alias is never applied to a different provider. + """ + async with create_session() as session: + query = select(UpstreamProviderRow).options( + selectinload(UpstreamProviderRow.models) # type: ignore[arg-type] + ) + provider_rows = (await session.exec(query)).all() + + overrides_by_key: dict[ModelKey, ModelRow] = {} + disabled_model_keys: set[ModelKey] = set() + provider_identities: dict[int, ConfiguredProviderIdentity] = {} + + for provider in provider_rows: + if not provider.enabled or provider.id is None: + continue + provider_identities[provider.id] = ConfiguredProviderIdentity( + id=provider.id, + slug=provider.slug or f"provider-{provider.id}", + provider_type=provider.provider_type, + base_url=public_provider_url(provider.base_url), + ) + for model in provider.models: + key = (model.id.lower(), provider.id) + if model.enabled: + overrides_by_key[key] = model + else: + disabled_model_keys.add(key) + + return overrides_by_key, disabled_model_keys, provider_identities + + +def _apply_model_visibility( + upstream: BaseUpstreamProvider, + overrides_by_key: dict[ModelKey, ModelRow] | None, + disabled_model_keys: set[ModelKey] | None, +) -> list[object]: + """Return provider models after DB disabled/override state is applied. + + Only the identity fields (``id``, ``forwarded_model_id``, + ``canonical_slug``) matter for path discovery, so DB override rows are used + directly rather than rebuilt into fully priced ``Model`` objects — the + pricing pipeline costs ~0.7ms of event-loop CPU per row for data this + module immediately discards. + """ + overrides_by_key = overrides_by_key or {} + disabled_model_keys = disabled_model_keys or set() + upstream_provider_id = getattr(upstream, "db_id", None) + if not isinstance(upstream_provider_id, int): + return [ + model + for model in upstream.get_cached_models() + if getattr(model, "enabled", True) + ] + + visible_models: list[object] = [] + seen_model_ids: set[str] = set() + + for model in upstream.get_cached_models(): + model_id = getattr(model, "id", "") + key = (model_id.lower(), upstream_provider_id) + if not getattr(model, "enabled", True) or key in disabled_model_keys: + continue + # Apply overrides only for this provider's own model row. + override_row = overrides_by_key.get(key) + visible: object = model if override_row is None else override_row + visible_models.append(visible) + seen_model_ids.add(model_id.lower()) + + # DB-only override rows for this provider with no cached counterpart. + for (model_id_lower, provider_id), override_row in overrides_by_key.items(): + if provider_id != upstream_provider_id: + continue + if model_id_lower in seen_model_ids: + continue + visible_models.append(override_row) + seen_model_ids.add(model_id_lower) + + return visible_models + + +async def _collect_provider_paths( + upstream: BaseUpstreamProvider, + provider_identity: ConfiguredProviderIdentity, + overrides_by_key: dict[ModelKey, ModelRow] | None = None, + disabled_model_keys: set[ModelKey] | None = None, + cycle: _RefreshCycleState | None = None, +) -> ProviderPathSnapshot: + """Collect selectable routes while marking model-level degraded fetches. + + A failed OpenRouter lookup preserves only that model's prior rows. Other + models in the same provider still refresh, so a partial outage cannot erase + valid discovery data or freeze the entire provider snapshot. + """ + cycle = cycle or _RefreshCycleState() + models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys) + + def _base_path(model: object) -> DiscoveredPath: + model_id = exposed_model_id(model) + return DiscoveredPath( + model_id=model_id, + path=encode_model_path( + provider_identity.base_url, provider_identity.id, model_id + ), + provider=provider_identity, + ) + + if not is_openrouter_base_url(upstream.base_url): + return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models)) + + if not (upstream.provider_type or "").strip(): + return ProviderPathSnapshot(paths=()) + + semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY) + async with _make_http_client() as client: + + async def _for_model( + model: object, + ) -> tuple[list[DiscoveredPath], str | None]: + model_id = exposed_model_id(model) + author_slug = openrouter_author_slug(model) + if not author_slug: + return [_base_path(model)], None + endpoints = await _fetch_openrouter_endpoint_subproviders( + client, + upstream.base_url, + upstream.api_key, + author_slug, + semaphore, + cycle, + ) + if endpoints is None: + return [], model_id + paths = [_base_path(model)] + paths.extend( + DiscoveredPath( + model_id=model_id, + path=encode_model_path( + provider_identity.base_url, + provider_identity.id, + model_id, + endpoint.tag, + ), + provider=provider_identity, + endpoint_tag=endpoint.tag, + endpoint_name=endpoint.provider_name, + ) + for endpoint in endpoints + ) + return paths, None + + results = await asyncio.gather( + *(_for_model(model) for model in models), return_exceptions=True + ) + + paths: list[DiscoveredPath] = [] + preserve_model_ids: set[str] = set() + for model, result in zip(models, results): + if isinstance(result, BaseException): + model_id = exposed_model_id(model) + preserve_model_ids.add(model_id) + logger.warning( + "OpenRouter endpoint discovery task errored", + extra={"provider": upstream.provider_type, "error": str(result)}, + ) + continue + model_paths, preserved_model_id = result + paths.extend(model_paths) + if preserved_model_id: + preserve_model_ids.add(preserved_model_id) + + return ProviderPathSnapshot( + paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids) + ) + + +async def _persist_provider_paths( + upstream_provider_id: int, snapshot: ProviderPathSnapshot +) -> None: + """Replace refreshed rows while retaining model-level degraded snapshots.""" + unique_paths = list( + {(path.model_id, path.path): path for path in snapshot.paths}.values() + ) + now = int(time.time()) + async with create_session() as session: + delete_stmt = delete(ModelPathRow).where( + col(ModelPathRow.upstream_provider_id) == upstream_provider_id + ) + if snapshot.preserve_model_ids: + delete_stmt = delete_stmt.where( + col(ModelPathRow.model_id).not_in(sorted(snapshot.preserve_model_ids)) + ) + await session.exec(delete_stmt) # type: ignore[call-overload] + for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE): + chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE] + values = [ + { + "model_id": discovered.model_id, + "path": discovered.path, + "provider_slug": discovered.provider.slug, + "provider_type": discovered.provider.provider_type, + "endpoint_tag": discovered.endpoint_tag, + "endpoint_name": discovered.endpoint_name, + "upstream_provider_id": upstream_provider_id, + "updated_at": now, + } + for discovered in chunk + ] + insert_stmt = insert(ModelPathRow).values(values) + await session.execute( + insert_stmt.on_conflict_do_update( + index_elements=["model_id", "path", "upstream_provider_id"], + set_={ + "provider_slug": insert_stmt.excluded.provider_slug, + "provider_type": insert_stmt.excluded.provider_type, + "endpoint_tag": insert_stmt.excluded.endpoint_tag, + "endpoint_name": insert_stmt.excluded.endpoint_name, + "updated_at": insert_stmt.excluded.updated_at, + }, + ) + ) + await session.commit() + + +async def prune_model_paths_for_inactive_providers() -> None: + """Delete paths whose provider is no longer enabled in the database. + + Called from ``refresh_model_maps`` so admin mutations (disable/delete + provider) stop advertising a provider's paths immediately instead of + waiting for the next timed refresh. Uses the DB as the source of truth, so + it is safe at boot even before upstreams initialize. + """ + async with create_session() as session: + enabled_ids = ( + await session.exec( + select(UpstreamProviderRow.id).where( + col(UpstreamProviderRow.enabled).is_(True) + ) + ) + ).all() + stmt = delete(ModelPathRow) + if enabled_ids: + stmt = stmt.where( + col(ModelPathRow.upstream_provider_id).not_in( + [pid for pid in enabled_ids if pid is not None] + ) + ) + await session.exec(stmt) # type: ignore[call-overload] + await session.commit() + + +async def refresh_model_paths( + upstreams: list[BaseUpstreamProvider], +) -> None: + """Recompute and persist model paths for every enabled provider. + + One provider's failure is logged and isolated; it must not break the rest. + A provider whose paths could not be determined this cycle keeps its + previously persisted rows. An empty ``upstreams`` list (e.g. a failed + ``initialize_upstreams`` at boot) is treated as "unknown" and touches + nothing. + """ + if not upstreams: + logger.warning("Skipping model paths refresh: no live upstreams") + return + + ( + overrides_by_key, + disabled_model_keys, + provider_identities, + ) = await _load_model_visibility() + await prune_model_paths_for_inactive_providers() + + cycle = _RefreshCycleState() + for upstream in upstreams: + if upstream.db_id is None or upstream.db_id not in provider_identities: + continue + try: + snapshot = await _collect_provider_paths( + upstream, + provider_identity=provider_identities[upstream.db_id], + overrides_by_key=overrides_by_key, + disabled_model_keys=disabled_model_keys, + cycle=cycle, + ) + if snapshot.preserve_model_ids: + logger.warning( + "Some model paths are unknown; keeping their previous rows", + extra={ + "provider": upstream.provider_type or upstream.base_url, + "db_id": upstream.db_id, + "preserved_models": len(snapshot.preserve_model_ids), + }, + ) + await _persist_provider_paths(upstream.db_id, snapshot) + except Exception as e: # noqa: BLE001 - isolate per-provider failures + logger.error( + "Failed to refresh model paths for provider", + extra={ + "provider": upstream.provider_type or upstream.base_url, + "db_id": upstream.db_id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + + +async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None: + """Synchronize one provider when model-path discovery is enabled.""" + if _refresh_interval_seconds() <= 0: + return + + from ..proxy import get_upstreams + + matching = [ + upstream + for upstream in get_upstreams() + if upstream.db_id == upstream_provider_id + ] + if matching: + await refresh_model_paths(matching) + else: + await prune_model_paths_for_inactive_providers() + + +async def _drain_scheduled_provider_refreshes() -> None: + """Serialize and coalesce model-path refreshes scheduled by admin writes.""" + global _scheduled_provider_refresh_task + + try: + # Let mutations in the same event-loop turn collapse into one refresh. + await asyncio.sleep(0) + while _scheduled_provider_refresh_ids: + if _refresh_interval_seconds() <= 0: + _scheduled_provider_refresh_ids.clear() + return + provider_id = min(_scheduled_provider_refresh_ids) + _scheduled_provider_refresh_ids.remove(provider_id) + try: + await refresh_model_paths_for_provider(provider_id) + except asyncio.CancelledError: + raise + except Exception as exc: # noqa: BLE001 - background best effort + logger.warning( + "Failed to refresh model paths after admin mutation", + extra={ + "upstream_provider_id": provider_id, + "error": str(exc), + "error_type": type(exc).__name__, + }, + ) + finally: + _scheduled_provider_refresh_task = None + + +async def schedule_model_paths_refresh_for_provider( + upstream_provider_id: int, +) -> None: + """Queue a non-blocking, coalesced refresh after an admin mutation.""" + global _scheduled_provider_refresh_task + + if _refresh_interval_seconds() <= 0: + return + _scheduled_provider_refresh_ids.add(upstream_provider_id) + if ( + _scheduled_provider_refresh_task is None + or _scheduled_provider_refresh_task.done() + ): + _scheduled_provider_refresh_task = asyncio.create_task( + _drain_scheduled_provider_refreshes(), + name="model-path-admin-refresh", + ) + + +def _refresh_interval_seconds() -> int: + """Current interval, re-read every loop so runtime setting changes apply.""" + from ..core.settings import settings + + if not getattr(settings, "enable_model_paths_refresh", True): + return 0 + return int(getattr(settings, "model_paths_refresh_interval_seconds", 0) or 0) + + +async def refresh_model_paths_periodically( + upstreams_provider: ( + Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider] + ), +) -> None: + """Background task mirroring ``refresh_upstreams_models_periodically``. + + The interval and enable flag are re-read every iteration, so the refresh + can be turned off (or on) and retuned without a restart. While disabled the + task idles instead of exiting, so re-enabling takes effect. + """ + _DISABLED_POLL_SECONDS = 60.0 + + def _resolve_upstreams() -> list[BaseUpstreamProvider]: + if callable(upstreams_provider): + return upstreams_provider() + return upstreams_provider + + while True: + interval = _refresh_interval_seconds() + if interval <= 0: + try: + await asyncio.sleep(_DISABLED_POLL_SECONDS) + except asyncio.CancelledError: + break + continue + + try: + await refresh_model_paths(_resolve_upstreams()) + except asyncio.CancelledError: + break + except Exception as e: # noqa: BLE001 + logger.error( + "Error in model paths refresh loop", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + try: + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) + except asyncio.CancelledError: + break + + +def _serialize_path(row: ModelPathRow) -> dict[str, Any]: + endpoint = None + if row.endpoint_tag or row.endpoint_name: + endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name} + return { + "path": row.path, + "provider": { + "id": row.upstream_provider_id, + "slug": row.provider_slug, + "type": row.provider_type, + }, + "endpoint": endpoint, + } + + +async def get_all_model_paths() -> dict: + """All models with their exact selectable routes.""" + async with create_session() as session: + rows = ( + await session.exec( + select(ModelPathRow).order_by( + col(ModelPathRow.model_id), + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) + ) + ).all() + + grouped: dict[str, list[dict[str, Any]]] = {} + seen_paths: dict[str, set[str]] = {} + updated_at = 0 + for row in rows: + updated_at = max(updated_at, row.updated_at) + if row.path in seen_paths.setdefault(row.model_id, set()): + continue + seen_paths[row.model_id].add(row.path) + grouped.setdefault(row.model_id, []).append(_serialize_path(row)) + data = [ + { + "id": grouped_model_id, + "paths": grouped[grouped_model_id], + } + for grouped_model_id in sorted(grouped) + ] + return {"data": data, "updated_at": updated_at or None} + + +async def get_paths_for_model(model_id: str) -> dict: + """Return paths for an advertised ID or its provider-prefixed alias.""" + + async def load_rows(session: AsyncSession, lookup_id: str) -> list[ModelPathRow]: + return list( + ( + await session.exec( + select(ModelPathRow) + .where(col(ModelPathRow.model_id) == lookup_id) + .order_by( + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) + ) + ).all() + ) + + async with create_session() as session: + rows = await load_rows(session, model_id) + if not rows: + unprefixed_id = public_model_id(model_id) + if unprefixed_id != model_id: + rows = await load_rows(session, unprefixed_id) + + seen: set[str] = set() + paths: list[dict] = [] + updated_at = 0 + for row in rows: + updated_at = max(updated_at, row.updated_at) + if row.path in seen: + continue + seen.add(row.path) + paths.append(_serialize_path(row)) + return {"data": paths, "updated_at": updated_at or None} diff --git a/routstr/wallet.py b/routstr/wallet.py index 8ff3c8f2..064080cd 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -2,7 +2,6 @@ import asyncio import fcntl import os import re -import socket import time import typing from contextlib import asynccontextmanager @@ -12,18 +11,37 @@ from pathlib import Path from typing import AsyncGenerator, TypedDict import httpx -from cashu.core.base import MeltQuote, Proof, Token +from cashu.core.base import MeltQuote, 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 +from cashu.wallet.wallet import Wallet as _CashuWallet from pydantic_core import PydanticUndefined 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. @@ -70,7 +88,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) @@ -100,6 +119,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). @@ -107,6 +140,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 @@ -118,15 +155,15 @@ class TokenConsumedError(Exception): """ -# httpx base classes cover their subclasses. HTTPStatusError is excluded on -# purpose — that means the mint answered, just with an error status. -_TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = ( - httpx.NetworkError, - httpx.TimeoutException, - ConnectionError, # refused/reset/aborted - socket.gaierror, # DNS failure - asyncio.TimeoutError, -) +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: @@ -144,7 +181,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 @@ -181,6 +218,20 @@ 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_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", @@ -259,58 +310,157 @@ 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]) + try: + await run_mint_operation( + lambda: wallet.load_mint(keyset_id=token_obj.keysets[0]), + op_name="redeem_load_mint", + mint_url=token_obj.mint, + ) + except Exception as error: + if is_mint_connection_error(error): + logger.warning( + "Same-mint redemption failed before swap dispatch", + 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_attempted": False, + "action": "retry_with_token_from_another_mint", + "error": str(error), + "error_type": type(error).__name__, + }, + ) + raise SourceMintConnectionError( + "Issuing Cashu mint is unreachable" + ) from error + raise + 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) + 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 destinations: + logger.info( + "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_trusted_mint( + token_obj, wallet, destination_mints=destinations + ) - if token_obj.mint not in settings.cashu_mints: - return await swap_to_primary_mint(token_obj, wallet) - + logger.info( + "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_on_connection_failure": False, + }, + ) return await _redeem_same_mint(wallet, token_obj) 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) - all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) - proofs = [proof for proof in all_proofs if not proof.reserved] - # Fallback must compare the requested amount with liquid proofs only. Counting - # reserved proofs here can suppress fallback even though they cannot be sent. - proofs_for_mint = sum(p.amount for p in proofs) - reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved) + """Create a token from the preferred mint or another funded trusted mint.""" + async with wallet_operation_guard(): + return await _send_locked(amount, unit, mint_url) - # 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) - all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) - proofs = [proof for proof in all_proofs if not proof.reserved] - proofs_for_mint = sum(p.amount for p in proofs) - reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved) + +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 + ) + proofs_for_mint = sum(proof.amount for proof in proofs) + all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) + reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved) 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 @@ -591,6 +741,74 @@ async def check_bolt11_payment_status(mint_url: str, unit: str, quote_id: str) - return "unknown" +async def release_token_reservation(token: str) -> None: + """Release a token that was created locally but never handed off.""" + 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: + 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, + *, + 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])) + 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, + force_reload=force_reload, + ) + 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 @@ -621,6 +839,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) @@ -657,13 +886,152 @@ 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, + 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. + + 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 = _trusted_destination_candidates(destination_mints) + logger.info( + "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 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.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.info( + "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 + else: + wallet = await get_wallet( + mint_url, + settings.primary_mint_unit, + retry_on_rate_limit=False, + ) + quote = await run_mint_operation( + lambda: wallet.request_mint(amount), + op_name=op_name, + mint_url=mint_url, + retry_on_rate_limit=False, + ) + logger.info( + "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, reason="unreachable" + ) + logger.warning( + "Destination mint failed", + extra={ + "event": "cashu_destination_failed", + "failed_mint": mint_url, + "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}") + + async def _calculate_swap_amount( amount_msat: int, token_unit: str, token_mint_url: str, token_wallet: Wallet, - primary_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 @@ -676,22 +1044,59 @@ 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) + # 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_trusted_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", + "swap_to_trusted_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, }, ) + stage = "destination_fee_quote" 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, + destination_mints=destination_mints, + ) + stage = "source_fee_quote" + 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, + ) fee_reserve = dummy_melt_quote.fee_reserve input_fees = token_wallet.get_fees_for_proofs(proofs) @@ -702,7 +1107,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, @@ -710,27 +1115,111 @@ 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 except Exception as e: logger.error( - "swap_to_primary_mint: fee estimation failed", - extra={"error": str(e)}, + "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, + "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): + 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, + "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 _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, + *, + destination_mints: list[str] | None = None, ) -> tuple[int, str, str]: logger.info( - "swap_to_primary_mint: starting", + "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, @@ -748,12 +1237,13 @@ async def swap_to_primary_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_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, @@ -762,7 +1252,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, @@ -771,6 +1261,7 @@ async def swap_to_primary_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 @@ -778,19 +1269,80 @@ 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) + if minted_amount <= 0: + logger.error( + "swap_to_trusted_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, + destination_mints=destination_candidates, + ) logger.info( - "swap_to_primary_mint: mint quote received", - extra={"mint_quote_id": mint_quote.quote, "attempt": attempt}, + "swap_to_trusted_mint: mint quote received", + extra={ + "mint_quote_id": mint_quote.quote, + "attempt": attempt, + "dest_mint": dest_mint_url, + }, ) - melt_quote = await token_wallet.melt_quote(mint_quote.request) + logger.info( + "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 run_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 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 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, @@ -810,7 +1362,7 @@ async def swap_to_primary_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, @@ -827,7 +1379,7 @@ async def swap_to_primary_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, @@ -839,32 +1391,56 @@ 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, + melt_response = await run_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, + 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( - "swap_to_primary_mint: melt failed — mint unreachable", - extra={"error": str(e), "foreign_mint": token_obj.mint}, - ) - raise MintConnectionError("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__, @@ -879,7 +1455,7 @@ async def swap_to_primary_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, @@ -892,34 +1468,46 @@ 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}, + "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, + }, ) - 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 run_mint_operation( + lambda: dest_wallet.mint(minted_amount, quote_id=mint_quote.quote), + op_name="swap_mint_on_destination", + 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(): # Previous mint call signed outputs at the mint but failed before # 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, }, ) 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", + "swap_to_trusted_mint: recovery scan completed", extra={ "pre_mint_balance": pre_mint_balance, "post_recovery_balance": post_recovery_balance, @@ -942,7 +1530,7 @@ async def swap_to_primary_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( @@ -950,7 +1538,7 @@ async def swap_to_primary_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__, @@ -964,17 +1552,25 @@ async def swap_to_primary_mint( ) from e logger.info( - "swap_to_primary_mint: completed successfully", + "Cashu cross-mint swap completed", extra={ - "foreign_mint": token_obj.mint, - "primary_mint": settings.primary_mint, + "event": "cashu_swap_completed", + "source_mint": token_obj.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 swap_to_primary_mint( + 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) async def credit_balance( @@ -988,12 +1584,22 @@ async def _credit_balance_locked( 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: - 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( @@ -1032,10 +1638,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 @@ -1089,18 +1704,51 @@ async def _credit_balance_locked( _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, +# 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 +async def get_wallet( + mint_url: str, + 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}" - 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: - await _wallets[id].load_mint() - await _wallets[id].load_proofs(reload=True) - return _wallets[id] + if load: + now = time.monotonic() + last = _wallet_last_load.get(id) + 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", + mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, + ) + await run_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] def get_proofs_per_mint_and_unit( @@ -1117,20 +1765,34 @@ 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 = [] _spent_proofs = [] - for i in range(0, len(proofs), 1000): - pb = proofs[i : i + 1000] - proof_states = await wallet.check_proof_state(pb) + # 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 run_mint_operation( + 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": _proofs.append(proof) else: _spent_proofs.append(proof) - await wallet.set_reserved_for_send(_spent_proofs, reserved=True) + if _spent_proofs: + await wallet.set_reserved_for_send(_spent_proofs, reserved=True) return _proofs @@ -1141,75 +1803,248 @@ 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 +_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]]] = {} + + +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 run_mint_operation( + lambda: wallet._get_keysets(), + op_name="get_mint_keysets", + mint_url=mint_url, + retry_on_rate_limit=False, + ) + 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: + 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( units: list[str] | None = None, ) -> tuple[list[BalanceDetail], int, int, int]: - """ - Fetch balances for all trusted mints and units concurrently. - - Returns: - - List of balance details for each mint/unit combination - - Total wallet balance in sats - - Total user balance in sats - - Owner balance in sats (wallet - user) - """ - if units is None: - units = ["sat", "msat"] - - # Received tokens are stored against primary_mint even when cashu_mints is - # empty, so include it in both the liability query and mint fan-out. + """Fetch balances for all trusted mints without holding DB connections during I/O.""" mint_urls = _mints_to_inspect() + mint_units: dict[str, list[str]] = {} + discovery_errors: list[BalanceDetail] = [] + if units is not None: + mint_units = {mint_url: units for mint_url in mint_urls} + else: + for mint_url in mint_urls: + try: + 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) + error_code = ( + "rate_limited" + if rate_limited + else "unreachable" + if connection_failure + else "mint_error" + ) + if connection_failure: + MintRateGuard.get(mint_url).apply_cooldown( + _BALANCE_FETCH_RETRY_SECONDS, reason="unreachable" + ) + retry_delay = max( + _BALANCE_FETCH_RETRY_SECONDS, + mint_cooldown_remaining(mint_url), + ) + discovery_errors.append( + _balance_error( + mint_url, + settings.primary_mint_unit, + str(error), + error_code=error_code, + retry_after_seconds=retry_delay, + ) + ) + mint_units[mint_url] = [] + 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__, + }, + ) + + # Read all liabilities in one short-lived transaction, then release the + # connection before starting concurrent mint network requests. user_balances: dict[tuple[str, str], int] = {} liabilities_error: str | None = None + query_units = list( + dict.fromkeys(unit for mint_url in mint_urls for unit in mint_units[mint_url]) + ) try: async with db.create_session() as session: user_balances = await db.balances_by_mint_and_unit( - session, mint_urls, units + session, mint_urls, query_units ) - except Exception as e: - logger.error("Error reading user balances", extra={"error": str(e)}) - liabilities_error = str(e) + except Exception as error: + logger.error("Error reading user balances", extra={"error": str(error)}) + liabilities_error = str(error) mint_check_limit = asyncio.Semaphore(settings.mint_operation_concurrency) async def fetch_balance(mint_url: str, unit: str) -> BalanceDetail: - try: - async with mint_check_limit: - wallet = await get_wallet(mint_url, unit) - proofs = get_proofs_per_mint_and_unit( - wallet, mint_url, unit, not_reserved=True + key = (mint_url, unit) + lock = _balance_fetch_locks.setdefault(mint_url, 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], + error_code=failure[2], + retry_after_seconds=failure[0] - now, ) - proofs = await slow_filter_spend_proofs(proofs, wallet) + + cooldown = mint_cooldown_remaining(mint_url) + if cooldown > 0: + 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: + async with mint_check_limit: + 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) + 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 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 + ) + retry_delay = max( + _BALANCE_FETCH_RETRY_SECONDS, + mint_cooldown_remaining(mint_url), + ) + _balance_fetch_failures[key] = ( + time.monotonic() + retry_delay, + str(error), + error_code, + ) + logger.warning( + "Unable to refresh mint balance", + extra={ + "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), + }, + ) + return _balance_error( + mint_url, + unit, + str(error), + error_code=error_code, + retry_after_seconds=retry_delay, + ) + + _balance_fetch_failures.pop(key, None) user_balance = user_balances.get((mint_url, unit), 0) if unit == "sat": user_balance = _msats_to_sats(user_balance) proofs_balance = sum(proof.amount for proof in proofs) - - result: BalanceDetail = { + return { "mint_url": mint_url, "unit": unit, "wallet_balance": proofs_balance, "user_balance": user_balance, "owner_balance": proofs_balance - user_balance, } - 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 - tasks = [fetch_balance(mint_url, unit) for mint_url in mint_urls for unit in units] - balance_details = list(await asyncio.gather(*tasks)) + tasks = [ + fetch_balance(mint_url, unit) + for mint_url in mint_urls + for unit in mint_units[mint_url] + ] + balance_details = discovery_errors + list(await asyncio.gather(*tasks)) total_wallet_balance_sats = 0 total_user_balance_sats = 0 @@ -1232,8 +2067,6 @@ async def fetch_all_balances( if liabilities_error is None: owner_balance = total_wallet_balance_sats - total_user_balance_sats else: - # Custody remains knowable when the DB read fails, but the user/owner - # split does not. Never report unknown liabilities as owner profit. owner_balance = 0 for detail in balance_details: detail["user_balance"] = 0 @@ -1251,7 +2084,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 ) @@ -1264,6 +2100,8 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: ) return + # Fetch liability after the proofs snapshot and settle delay while the + # wallet operation guard excludes concurrent proof mutation and crediting. try: async with db.create_session() as session: # ApiKey stores a refund preference, not funding provenance. Until @@ -1322,8 +2160,7 @@ async def periodic_payout() -> None: for mint_url in _mints_to_inspect(): for unit in ["sat", "msat"]: # Proof mutation, liability observation, and sending are one - # cross-process critical section. A credit takes the same - # lock from before redemption through its liability commit. + # cross-process critical section. Credits take the same lock. async with wallet_operation_guard(): await _payout_mint_and_unit(mint_url, unit) except Exception as e: @@ -1389,7 +2226,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, @@ -1520,66 +2358,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__}", @@ -1588,10 +2433,18 @@ 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) - 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_insufficient_balance.py b/tests/integration/test_insufficient_balance.py index fdc04e2c..a63c4251 100644 --- a/tests/integration/test_insufficient_balance.py +++ b/tests/integration/test_insufficient_balance.py @@ -207,8 +207,35 @@ async def test_pay_for_request_succeeds_when_balance_equals_cost( assert key.balance == model_cost # balance unchanged, only reserved goes up +@pytest.mark.asyncio +async def test_full_model_maximum_is_required_and_reserved( + integration_session: AsyncSession, +) -> None: + from routstr.auth import pay_for_request, validate_bearer_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() + + 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 + + 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 + + # --------------------------------------------------------------------------- -# Test 6 — HTTP layer returns 402 JSON with the right shape +# HTTP layer returns 402 JSON with the right shape # --------------------------------------------------------------------------- @pytest.mark.asyncio @@ -266,8 +293,8 @@ 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 str(model_cost) in error["message"] - assert str(user_balance) 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 await integration_session.refresh(key) diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index 079954ce..1a6d94b9 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -14,6 +14,7 @@ import time from unittest.mock import AsyncMock, MagicMock, patch import pytest +from cashu.core.base import Proof from sqlalchemy import inspect from sqlalchemy.ext.asyncio import AsyncEngine from sqlmodel.ext.asyncio.session import AsyncSession @@ -22,6 +23,12 @@ from routstr.core.db import ApiKey, LightningInvoice from routstr.lightning import _create_api_key_record +def _configure_quote_proof_wallet(wallet: MagicMock) -> None: + wallet.proofs = [] + wallet.keysets = {} + wallet.load_proofs = AsyncMock() + + def _make_invoice(**kwargs: object) -> LightningInvoice: base = dict( id="inv_test_001", @@ -42,7 +49,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 @@ -183,6 +198,7 @@ async def test_concurrent_payment_checks_mint_and_credit_invoice_once( await setup.commit() wallet = MagicMock() + _configure_quote_proof_wallet(wallet) wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) mint_calls = 0 @@ -196,7 +212,9 @@ async def test_concurrent_payment_checks_mint_and_credit_invoice_once( await asyncio.sleep(0.05) if call_number > 1: raise Exception("quote already issued") - return [] + proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash) + wallet.proofs.append(proof) + return [proof] wallet.mint = AsyncMock(side_effect=single_use_mint) @@ -228,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: @@ -238,6 +256,7 @@ async def test_failed_mint_keeps_invoice_pending_for_retry( await setup.commit() wallet = MagicMock() + _configure_quote_proof_wallet(wallet) wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) wallet.mint = AsyncMock(side_effect=TimeoutError("mint unavailable")) async with AsyncSession(integration_engine, expire_on_commit=False) as session: @@ -251,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 @@ -349,8 +368,15 @@ async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation( await setup.commit() wallet = MagicMock() + _configure_quote_proof_wallet(wallet) wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) - wallet.mint = AsyncMock(return_value=[]) + + async def successful_mint(*args: object, **kwargs: object) -> list[Proof]: + proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash) + wallet.proofs.append(proof) + return [proof] + + wallet.mint = AsyncMock(side_effect=successful_mint) async with AsyncSession(integration_engine, expire_on_commit=False) as session: stored = await session.get(LightningInvoice, invoice.id) stored_sibling = await session.get(LightningInvoice, sibling.id) @@ -373,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 @@ -428,11 +454,14 @@ async def test_db_guard_credits_once_when_both_mints_succeed( await setup.commit() wallet = MagicMock() + _configure_quote_proof_wallet(wallet) wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) - async def always_succeeding_mint(*args: object, **kwargs: object) -> list[object]: + async def always_succeeding_mint(*args: object, **kwargs: object) -> list[Proof]: await asyncio.sleep(0.05) - return [] + proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash) + wallet.proofs.append(proof) + return [proof] wallet.mint = AsyncMock(side_effect=always_succeeding_mint) @@ -472,7 +501,7 @@ async def test_db_guard_credits_once_when_both_mints_succeed( assert first_invoice not in first.dirty assert second_invoice not in second.dirty - assert wallet.mint.await_count == 2 + assert wallet.mint.await_count == 1 async with AsyncSession(integration_engine, expire_on_commit=False) as verify: stored_invoice = await verify.get(LightningInvoice, invoice.id) assert stored_invoice is not None diff --git a/tests/integration/test_lightning_invoice_rip08.py b/tests/integration/test_lightning_invoice_rip08.py index 29301a42..faba77c7 100644 --- a/tests/integration/test_lightning_invoice_rip08.py +++ b/tests/integration/test_lightning_invoice_rip08.py @@ -26,11 +26,17 @@ 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, + *, + allowed_mints: list[str] | None = None, + ) -> 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( @@ -95,6 +101,8 @@ async def test_topup_with_authorization_header( body = resp.json() assert body["amount_sats"] == 500 assert body["bolt11"].startswith("lnbc") + allowed_mints = patch_invoice_generation.call_args.kwargs["allowed_mints"] + assert allowed_mints == ["http://localhost:3338"] @pytest.mark.integration diff --git a/tests/integration/test_lightning_settlement.py b/tests/integration/test_lightning_settlement.py new file mode 100644 index 00000000..f38d6623 --- /dev/null +++ b/tests/integration/test_lightning_settlement.py @@ -0,0 +1,365 @@ +import asyncio +import time +import uuid +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 + +from routstr.core.db import ApiKey, LightningInvoice +from routstr.lightning import ( + _expire_invoice_if_authoritatively_unpaid, + _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: AsyncEngine, +) -> 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: AsyncEngine, +) -> 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: AsyncEngine, +) -> 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: AsyncEngine, + patched_db_engine: None, +) -> 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 + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch( + "routstr.lightning._finalize_invoice_settlement", + AsyncMock(side_effect=Exception("db unavailable")), + ), + ): + 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 == "settlement_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() + + +@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 d2326a89..8195a7a8 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 @@ -28,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( @@ -81,7 +84,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 @@ -89,7 +94,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 @@ -139,7 +149,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_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_auto_topup.py b/tests/unit/test_auto_topup.py index 8d9d5a57..f74af477 100644 --- a/tests/unit/test_auto_topup.py +++ b/tests/unit/test_auto_topup.py @@ -89,6 +89,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()) @@ -97,7 +101,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", @@ -161,8 +165,14 @@ async def test_auto_topup_does_not_send_untracked_token() -> None: "routstr.upstream.auto_topup.store_cashu_transaction", AsyncMock(side_effect=RuntimeError("database unavailable")), ), + 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_balance.py b/tests/unit/test_balance.py index f3775e22..1bf94c94 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -606,6 +606,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 @@ -758,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_cost_calculation_caching.py b/tests/unit/test_cost_calculation_caching.py index ba31366a..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,11 +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 # @@ -568,12 +694,14 @@ 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 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_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_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_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index 8eae9829..17be72ec 100644 --- a/tests/unit/test_fee_payout_migration.py +++ b/tests/unit/test_fee_payout_migration.py @@ -4,6 +4,9 @@ import subprocess import sys from pathlib import Path +from alembic.config import Config +from alembic.script import ScriptDirectory + def _run_alembic(root: Path, database_url: str, revision: str) -> None: env = os.environ.copy() @@ -37,7 +40,10 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None: "payout_in_progress_msats, payout_started_at FROM routstr_fees" ).fetchone() - assert version == ("aa50fde387a2",) + 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_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index 3427cf78..ceeca1e4 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -1,9 +1,10 @@ import asyncio -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Generator from contextlib import asynccontextmanager from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from sqlalchemy.ext.asyncio import create_async_engine from sqlmodel import SQLModel @@ -12,6 +13,21 @@ from sqlmodel.ext.asyncio.session import AsyncSession 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() + 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() + + @asynccontextmanager async def _fake_session(): # type: ignore[no-untyped-def] yield MagicMock() @@ -29,7 +45,7 @@ def _patches( # 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_by_mint_and_unit", @@ -63,6 +79,161 @@ 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 + + # 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() + 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 + + 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.mint.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 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() + + 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.mint.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_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 + 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.mint.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[0]["error_code"] == "unreachable" + assert details[1]["error"] == "Mint is unreachable" + assert details[1]["error_code"] == "unreachable" + + @pytest.mark.asyncio async def test_fetch_all_balances_closes_db_session_before_concurrent_mint_io() -> None: """Slow mint checks must never run while the balance DB session is open.""" diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py new file mode 100644 index 00000000..584950d6 --- /dev/null +++ b/tests/unit/test_lightning_settlement.py @@ -0,0 +1,370 @@ +import asyncio +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock, patch + +import httpx +import pytest +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 + + +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, + "created_at": 1, + "expires_at": 2, + } + 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() + assert _invoice_settlement_locks == {} + + +@pytest.mark.asyncio +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")), + ), + patch("routstr.lightning._reload_invoice_view", AsyncMock()), + ): + await check_invoice_payment(invoice, session) # type: ignore[arg-type] + + 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() + 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) + + @asynccontextmanager + async def owned_session() -> AsyncIterator[AsyncMock]: + owned = AsyncMock() + owned.exec.return_value.rowcount = 1 + yield owned + + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning.create_session", owned_session), + 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() + assert _invoice_settlement_locks == {} diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index 47efbcf2..b5567570 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -1,70 +1,202 @@ -"""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 httpx import pytest +from cashu.core.base import MeltQuoteState -from routstr.payment import lnurl -from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl +from routstr.core.settings import settings +from routstr.mint import MintCooldownError, MintRateGuard +from routstr.payment.lnurl import ( + MeltOutcomeAmbiguousError, + raw_send_to_lnurl, +) + +LNURL_DATA = { + "callback_url": "https://ln.tld/cb", + "min_sendable": 1_000, + "max_sendable": 100_000_000, +} + + +def _wallet() -> tuple[MagicMock, list[MagicMock]]: + proofs = [MagicMock(amount=1000)] + 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_times_out_on_hung_melt() -> None: - proofs = [MagicMock(amount=1000)] - - wallet = MagicMock() - wallet.melt_quote = AsyncMock(return_value=MagicMock(fee_reserve=1, quote="q")) - wallet.select_to_send = AsyncMock(return_value=(proofs, None)) +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(MeltOutcomeAmbiguousError, 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(MeltOutcomeAmbiguousError, 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 +@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() + 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 +204,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_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index c41568ad..cb3fe9d7 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 @@ -854,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 @@ -872,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 # --------------------------------------------------------------------------- diff --git a/tests/unit/test_mint.py b/tests/unit/test_mint.py new file mode 100644 index 00000000..6b70a558 --- /dev/null +++ b/tests/unit/test_mint.py @@ -0,0 +1,121 @@ +import asyncio +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_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: + 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) + + +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_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_mint_url_migration.py b/tests/unit/test_mint_url_migration.py new file mode 100644 index 00000000..b83566f1 --- /dev/null +++ b/tests/unit/test_mint_url_migration.py @@ -0,0 +1,47 @@ +import os +import sqlite3 +import subprocess +import sys +from pathlib import Path + + +def _run_alembic(root: Path, database_url: str, command: str, revision: str) -> None: + env = os.environ.copy() + env["DATABASE_URL"] = database_url + subprocess.run( + [sys.executable, "-m", "alembic", command, revision], + cwd=root, + env=env, + check=True, + capture_output=True, + text=True, + ) + + +def _lightning_invoice_columns(database_path: Path) -> set[str]: + with sqlite3.connect(database_path) as connection: + return { + row[1] + for row in connection.execute("PRAGMA table_info(lightning_invoices)") + } + + +def test_mint_url_migration_upgrades_and_downgrades_from_main_head( + tmp_path: Path, +) -> None: + root = Path(__file__).resolve().parents[2] + database_path = tmp_path / "mint-url-migration.db" + database_url = f"sqlite+aiosqlite:///{database_path}" + 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", "ecfa0d6e2a36") + assert "mint_url" in _lightning_invoice_columns(database_path) + + _run_alembic(root, database_url, "downgrade", previous_head) + assert "mint_url" not in _lightning_invoice_columns(database_path) + + _run_alembic(root, database_url, "upgrade", "head") + assert "mint_url" in _lightning_invoice_columns(database_path) diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py new file mode 100644 index 00000000..240a4199 --- /dev/null +++ b/tests/unit/test_model_paths.py @@ -0,0 +1,1475 @@ +"""Tests for the model-path discovery service and endpoints. + +These tests exercise the public entry points (``refresh_model_paths``, +``get_all_model_paths``, ``get_paths_for_model``) rather than private helpers, +and fake OpenRouter at the transport level (``httpx.MockTransport``) so a +signature drift in the SUT fails loudly instead of silently returning ``[]``. +""" + +from __future__ import annotations + +import asyncio +import json +import os +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Any, AsyncGenerator, Callable, cast + +import httpx +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from sqlalchemy import event +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.db import ModelRow, UpstreamProviderRow # noqa: E402 +from routstr.payment.models import models_router # noqa: E402 +from routstr.upstream import model_paths as mp # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 +from routstr.upstream.openrouter import OpenRouterUpstreamProvider # noqa: E402 + +# --------------------------------------------------------------------------- # +# Fakes +# --------------------------------------------------------------------------- # + + +def _model( + id: str, + *, + forwarded_model_id: str | None = None, + canonical_slug: str | None = None, + enabled: bool = True, +) -> SimpleNamespace: + return SimpleNamespace( + id=id, + forwarded_model_id=forwarded_model_id, + canonical_slug=canonical_slug, + enabled=enabled, + ) + + +def _model_row( + id: str, + *, + upstream_provider_id: int = 1, + forwarded_model_id: str | None = None, + canonical_slug: str | None = None, + enabled: bool = True, +) -> ModelRow: + return ModelRow( + id=id, + upstream_provider_id=upstream_provider_id, + name=id, + created=0, + description="test model", + context_length=8192, + architecture=json.dumps( + { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "test", + "instruct_type": None, + } + ), + pricing=json.dumps({"prompt": 0.000001, "completion": 0.000002}), + enabled=enabled, + forwarded_model_id=forwarded_model_id, + canonical_slug=canonical_slug, + ) + + +class _FakeProvider(BaseUpstreamProvider): + """Real ``BaseUpstreamProvider`` so the discovery-path hooks are the + production ones, with cached models injected.""" + + def __init__( + self, + *, + provider_type: str, + base_url: str, + models: list[SimpleNamespace], + db_id: int | None = 1, + api_key: str = "sk-test", + ) -> None: + super().__init__(base_url=base_url, api_key=api_key) + self.provider_type = provider_type # shadow the class attribute + self.db_id = db_id + self._models = models + + def get_cached_models(self) -> list[SimpleNamespace]: # type: ignore[override] + return self._models + + +class _FakeOpenRouterProvider(OpenRouterUpstreamProvider): + """Real OpenRouter provider so the ``unknown`` mapping is the production one.""" + + def __init__( + self, + *, + models: list[SimpleNamespace], + db_id: int | None = 2, + api_key: str = "sk-or", + ) -> None: + super().__init__(api_key=api_key) + self.db_id = db_id + self._models = models + + def get_cached_models(self) -> list[SimpleNamespace]: # type: ignore[override] + return self._models + + +def _mock_transport( + monkeypatch: pytest.MonkeyPatch, + handler: Callable[[httpx.Request], httpx.Response], +) -> dict[str, int]: + """Route the SUT's HTTP through ``httpx.MockTransport`` and count requests.""" + counter = {"requests": 0} + + def _counting_handler(request: httpx.Request) -> httpx.Response: + counter["requests"] += 1 + return handler(request) + + def _factory() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(_counting_handler)) + + monkeypatch.setattr(mp, "_make_http_client", _factory) + return counter + + +def _endpoints_response( + *providers: str | tuple[str, str], +) -> httpx.Response: + endpoints = [] + for provider in providers: + if isinstance(provider, tuple): + provider_name, tag = provider + else: + provider_name = provider + tag = provider.lower().replace(" ", "-") + endpoints.append({"provider_name": provider_name, "tag": tag}) + return httpx.Response(200, json={"data": {"endpoints": endpoints}}) + + +_SEEDED_PROVIDER_IDS = (1, 2, 4, 5, 7) + + +@pytest.fixture +async def patched_session( + monkeypatch: pytest.MonkeyPatch, +) -> AsyncGenerator[AsyncEngine, None]: + """Bind the service's ``create_session`` to a fresh in-memory engine. + + Foreign keys are enforced (``PRAGMA foreign_keys=ON``) so a ModelPathRow + insert for an unseeded provider fails here even though production SQLite + currently runs with the pragma off. + """ + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + + @event.listens_for(engine.sync_engine, "connect") + def _enable_fk(dbapi_conn: Any, _record: Any) -> None: + dbapi_conn.execute("PRAGMA foreign_keys=ON") + + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + # Seed every provider id the tests insert path rows for. + async with AsyncSession(engine) as session: + for pid in _SEEDED_PROVIDER_IDS: + session.add( + UpstreamProviderRow( + id=pid, + slug=f"p{pid}", + provider_type="anthropic" if pid == 1 else "openrouter", + base_url=f"https://provider-{pid}", + api_key=f"k{pid}", + ) + ) + await session.commit() + + @asynccontextmanager + async def _factory() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + monkeypatch.setattr(mp, "create_session", _factory) + yield engine + await engine.dispose() + + +def _paths_of(payload: dict, model_id: str) -> set[str]: + for entry in payload["data"]: + if entry["id"] == model_id: + return {p["path"] for p in entry["paths"]} + return set() + + +def _ids_of(payload: dict) -> set[str]: + return {entry["id"] for entry in payload["data"]} + + +def _expected_path( + provider_id: int, + model_id: str, + endpoint_tag: str | None = None, +) -> str: + return mp.encode_model_path( + f"https://provider-{provider_id}", provider_id, model_id, endpoint_tag + ) + + +def _path_entry( + provider_id: int, + model_id: str, + *, + provider_slug: str | None = None, + provider_type: str | None = None, + endpoint_tag: str | None = None, + endpoint_name: str | None = None, +) -> dict[str, Any]: + endpoint = None + if endpoint_tag or endpoint_name: + endpoint = {"tag": endpoint_tag, "name": endpoint_name} + return { + "path": _expected_path(provider_id, model_id, endpoint_tag), + "provider": { + "id": provider_id, + "slug": provider_slug or f"p{provider_id}", + "type": provider_type + or ("anthropic" if provider_id == 1 else "openrouter"), + }, + "endpoint": endpoint, + } + + +# --------------------------------------------------------------------------- # +# Predicates / pure helpers +# --------------------------------------------------------------------------- # + + +def test_is_openrouter_base_url() -> None: + assert mp.is_openrouter_base_url("https://openrouter.ai/api/v1") is True + assert mp.is_openrouter_base_url("https://api.anthropic.com") is False + assert mp.is_openrouter_base_url(None) is False + + +def test_native_anthropic_not_openrouter() -> None: + """Native Anthropic must not be treated as OpenRouter-compatible even though + ``_upstream_accepts_cache_control`` returns True for it.""" + assert mp.is_openrouter_base_url("https://api.anthropic.com/v1") is False + + +def test_public_provider_url_masks_private_addresses_and_explicit_ports() -> None: + assert mp.public_provider_url("http://192.168.1.10/v1") == "http://localhost" + assert mp.public_provider_url("http://10.0.0.5:11434/v1") == "http://localhost" + assert mp.public_provider_url("http://[fd00::1]/v1") == "http://localhost" + assert mp.public_provider_url("https://api.example.com:8443/v1") == ( + "http://localhost" + ) + + +def test_public_provider_url_preserves_public_urls_without_ports() -> None: + assert mp.public_provider_url("https://openrouter.ai/api/v1") == ( + "https://openrouter.ai/api/v1" + ) + assert mp.public_provider_url("http://localhost") == "http://localhost" + + +def test_encode_model_path_includes_complete_route_identity() -> None: + assert mp.encode_model_path( + "https://openrouter.ai/api/v1", 42, "anthropic/claude-sonnet-4" + ) == ( + "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1" + "&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4" + ) + assert mp.encode_model_path( + "https://openrouter.ai/api/v1", + 42, + "anthropic/claude-sonnet-4", + "google-vertex/us-east5", + ) == ( + "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1" + "&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4" + "&endpoint=google-vertex%2Fus-east5" + ) + + +def test_exposed_model_id_prefers_forwarded() -> None: + assert ( + mp.exposed_model_id(_model("claude-x", forwarded_model_id="fwd-claude")) + == "fwd-claude" + ) + assert mp.exposed_model_id(_model("claude-x")) == "claude-x" + + +def test_public_model_id_strips_first_provider_prefix() -> None: + """Must match ``create_model_mappings.get_base_model_id`` (first slash), + so the id shown by discovery can be sent to chat completions verbatim.""" + assert mp.public_model_id("z-ai/glm-5v-turbo") == "glm-5v-turbo" + assert mp.public_model_id("gpt-4o-mini") == "gpt-4o-mini" + assert ( + mp.public_model_id("accounts/fireworks/models/glm-5") + == "fireworks/models/glm-5" + ) + + +def test_openrouter_author_slug_prefers_canonical() -> None: + m = _model( + "claude-opus-4.6", + forwarded_model_id="forwarded-only", + canonical_slug="anthropic/claude-opus-4.6", + ) + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_falls_back_to_slash_id() -> None: + m = _model("anthropic/claude-opus-4.6", canonical_slug="claude-opus-4.6") + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_falls_back_to_forwarded_id() -> None: + """Admin-created alias rows have a slash-less local id; the forwarded id is + what the proxy actually sends to OpenRouter, so it is a usable slug.""" + m = _model("my-alias", forwarded_model_id="anthropic/claude-opus-4.6") + assert mp.openrouter_author_slug(m) == "anthropic/claude-opus-4.6" + + +def test_openrouter_author_slug_none_when_no_slash() -> None: + m = _model("claude-opus-4.6", canonical_slug="claude-opus-4.6") + assert mp.openrouter_author_slug(m) is None + + +# --------------------------------------------------------------------------- # +# Refresh through the public entry point +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_direct_provider_single_path_uses_provider_type( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + payload = await mp.get_all_model_paths() + assert payload["data"] == [ + {"id": "claude-opus-4.6", "paths": [_path_entry(1, "claude-opus-4.6")]} + ] + assert payload["updated_at"] is not None + + +@pytest.mark.asyncio +async def test_direct_path_masks_private_configured_provider_url( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.base_url = "http://192.168.1.10:11434/v1" + session.add(provider_row) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="http://192.168.1.10:11434/v1", + models=[_model("local-model")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "local-model") == { + mp.encode_model_path("http://localhost", 1, "local-model") + } + + +@pytest.mark.asyncio +async def test_direct_path_stores_exposed_model_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("internal-id", forwarded_model_id="claude-opus-4.6")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} + + +@pytest.mark.asyncio +async def test_forwarded_model_id_with_slash_remains_exact_and_routable( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("local-alias", forwarded_model_id="anthropic/claude-opus-4.6")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _ids_of(await mp.get_all_model_paths()) == {"anthropic/claude-opus-4.6"} + assert (await mp.get_paths_for_model("anthropic/claude-opus-4.6"))["data"] + + +@pytest.mark.asyncio +async def test_disabled_cached_models_excluded( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("enabled-model"), _model("disabled-model", enabled=False)], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"enabled-model"} + + +@pytest.mark.asyncio +async def test_disabling_model_on_one_provider_keeps_other_provider( + patched_session: AsyncEngine, +) -> None: + """Regression for cross-provider isolation: ModelRow's primary key is + (id, upstream_provider_id), so a disable row on provider 2 must not hide + provider 1's model.""" + async with AsyncSession(patched_session) as session: + session.add(_model_row("shared-model", upstream_provider_id=2, enabled=False)) + await session.commit() + + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("shared-model")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("shared-model")], + db_id=2, + ) + + await mp.refresh_model_paths([p1, p2]) + + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} + + +@pytest.mark.asyncio +async def test_override_alias_not_applied_across_providers( + patched_session: AsyncEngine, +) -> None: + """Provider 2's forwarded_model_id must never rename provider 1's model.""" + async with AsyncSession(patched_session) as session: + session.add( + _model_row( + "shared-model", + upstream_provider_id=2, + forwarded_model_id="private-alias", + ) + ) + await session.commit() + + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("shared-model")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("shared-model")], + db_id=2, + ) + + await mp.refresh_model_paths([p1, p2]) + + payload = await mp.get_all_model_paths() + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} + assert _paths_of(payload, "private-alias") == {_expected_path(2, "private-alias")} + + +@pytest.mark.asyncio +async def test_override_matching_is_case_insensitive( + patched_session: AsyncEngine, +) -> None: + """Routing lowercases both sides when matching DB rows to cached models; + discovery must do the same for mixed-case ids.""" + async with AsyncSession(patched_session) as session: + session.add( + _model_row( + "deepseek-ai/deepseek-v4-flash", + upstream_provider_id=1, + forwarded_model_id="public-alias", + ) + ) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("deepseek-ai/DeepSeek-V4-Flash")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"public-alias"} + + +@pytest.mark.asyncio +async def test_refresh_model_paths_excludes_db_disabled_override( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("disabled-by-db", enabled=False)) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("disabled-by-db")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + + assert (await mp.get_all_model_paths())["data"] == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_uses_db_forwarded_alias( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("internal-id", forwarded_model_id="public-alias")) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("internal-id")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + + assert (await mp.get_all_model_paths())["data"] == [ + {"id": "public-alias", "paths": [_path_entry(1, "public-alias")]} + ] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_includes_enabled_db_override_missing_from_cache( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + session.add(_model_row("deployment-id", forwarded_model_id="public-deployment")) + await session.commit() + + provider = _FakeProvider( + provider_type="generic", + base_url="https://custom-provider/v1", + models=[], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + + assert (await mp.get_all_model_paths())["data"] == [ + { + "id": "public-deployment", + "paths": [_path_entry(1, "public-deployment")], + } + ] + + +@pytest.mark.asyncio +async def test_refresh_replaces_stale_rows( + patched_session: AsyncEngine, +) -> None: + p_two_models = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1"), _model("m2")], + db_id=1, + ) + await mp.refresh_model_paths([p_two_models]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1", "m2"} + + p_one_model = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([p_one_model]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + +@pytest.mark.asyncio +async def test_refresh_with_no_upstreams_keeps_existing_rows( + patched_session: AsyncEngine, +) -> None: + """An empty live upstream list (e.g. failed boot init) means "unknown", + not "delete everything".""" + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + await mp.refresh_model_paths([]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + +@pytest.mark.asyncio +async def test_prune_removes_rows_of_disabled_db_provider( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("m1")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + assert _ids_of(await mp.get_all_model_paths()) == {"m1"} + + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.enabled = False + session.add(provider_row) + await session.commit() + + await mp.prune_model_paths_for_inactive_providers() + assert (await mp.get_all_model_paths())["data"] == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_skips_disabled_db_provider( + patched_session: AsyncEngine, +) -> None: + async with AsyncSession(patched_session) as session: + provider_row = await session.get(UpstreamProviderRow, 1) + assert provider_row is not None + provider_row.enabled = False + session.add(provider_row) + await session.commit() + + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("fresh-model")], + db_id=1, + ) + + await mp.refresh_model_paths([provider]) + + assert (await mp.get_all_model_paths())["data"] == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_skips_provider_without_db_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=None, + ) + await mp.refresh_model_paths([provider]) + assert (await mp.get_all_model_paths())["data"] == [] + + +@pytest.mark.asyncio +async def test_refresh_model_paths_isolates_provider_failure( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + good = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + bad = _FakeProvider( + provider_type="generic", + base_url="https://other-upstream/v1", + models=[_model("m")], + db_id=2, + ) + + original = mp._collect_provider_paths + + async def _maybe_fail(upstream: Any, *args: Any, **kwargs: Any) -> Any: + if upstream is bad: + raise RuntimeError("boom") + return await original(upstream, *args, **kwargs) + + monkeypatch.setattr(mp, "_collect_provider_paths", _maybe_fail) + + await mp.refresh_model_paths([good, bad]) + assert _ids_of(await mp.get_all_model_paths()) == {"claude-opus-4.6"} + + +# --------------------------------------------------------------------------- # +# OpenRouter endpoint discovery (transport-level fakes) +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_openrouter_provider_adds_endpoint_paths( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport( + monkeypatch, + lambda request: _endpoints_response( + ("Google", "google-vertex/eu"), + ("Google", "google-vertex/us"), + ), + ) + + await mp.refresh_model_paths([provider]) + + payload = await mp.get_paths_for_model("claude-opus-4.6") + assert {item["path"] for item in payload["data"]} == { + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "google-vertex/eu"), + _expected_path(2, "claude-opus-4.6", "google-vertex/us"), + } + assert { + item["endpoint"]["tag"] for item in payload["data"] if item["endpoint"] + } == {"google-vertex/eu", "google-vertex/us"} + assert { + item["endpoint"]["name"] for item in payload["data"] if item["endpoint"] + } == {"Google"} + assert {item["provider"]["id"] for item in payload["data"]} == {2} + + +@pytest.mark.asyncio +async def test_openrouter_uses_exact_tag_even_when_display_name_is_router( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Machine-readable endpoint tags, not display names, define identity.""" + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("OpenRouter")) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert paths == { + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "openrouter"), + } + + +@pytest.mark.asyncio +async def test_generic_provider_with_openrouter_base_url_discovers( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Configured provider identity is independent from its endpoint URL.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=1, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + + await mp.refresh_model_paths([provider]) + + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert paths == { + _expected_path(1, "claude-opus-4.6"), + _expected_path(1, "claude-opus-4.6", "anthropic"), + } + + +@pytest.mark.asyncio +async def test_openrouter_partial_failure_keeps_failed_models_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + provider = _FakeOpenRouterProvider( + models=[ + _model("good", canonical_slug="author/good"), + _model("degraded", canonical_slug="author/degraded"), + ], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "degraded") + assert before + + def _partial_failure(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/author/degraded/endpoints"): + return httpx.Response(503) + return _endpoints_response("Google") + + _mock_transport(monkeypatch, _partial_failure) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "degraded") == before + assert _paths_of(await mp.get_all_model_paths(), "good") != before + + +@pytest.mark.asyncio +async def test_partial_failure_upserts_collapsed_public_model_paths( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A degraded canonical sibling may preserve the same public path that a + successful sibling refreshes; persistence must merge instead of rolling back.""" + provider = _FakeOpenRouterProvider( + models=[ + _model("vendora/shared", canonical_slug="vendora/shared"), + _model("vendorb/shared", canonical_slug="vendorb/shared"), + ], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + + def _one_sibling_degrades(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/vendorb/shared/endpoints"): + return httpx.Response(503) + return _endpoints_response("Google") + + _mock_transport(monkeypatch, _one_sibling_degrades) + await mp.refresh_model_paths([provider]) + + assert _paths_of(await mp.get_all_model_paths(), "shared") == { + _expected_path(2, "shared"), + _expected_path(2, "shared", "anthropic"), + _expected_path(2, "shared", "google"), + } + + +@pytest.mark.asyncio +async def test_openrouter_failure_keeps_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """A transient upstream failure means the path set is unknown; previously + persisted rows must survive, mirroring ``refresh_models_cache``.""" + provider = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert _expected_path(2, "claude-opus-4.6", "anthropic") in before + + def _network_down(request: httpx.Request) -> httpx.Response: + raise httpx.ConnectError("network down", request=request) + + _mock_transport(monkeypatch, _network_down) + await mp.refresh_model_paths([provider]) + + after = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert after == before + + +@pytest.mark.asyncio +async def test_openrouter_rate_limit_aborts_cycle_and_keeps_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """The first 429 latches: no further endpoint requests this cycle, and the + provider's previously persisted rows survive.""" + models = [_model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(10)] + provider = _FakeOpenRouterProvider(models=models, db_id=2) + + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + expected = { + _expected_path(2, "m0"), + _expected_path(2, "m0", "anthropic"), + } + assert _paths_of(await mp.get_all_model_paths(), "m0") == expected + + counter = _mock_transport(monkeypatch, lambda request: httpx.Response(429)) + await mp.refresh_model_paths([provider]) + + # Up to _OPENROUTER_CONCURRENCY requests may already be in flight when the + # first 429 lands; the latch must stop everything after that. + assert counter["requests"] <= mp._OPENROUTER_CONCURRENCY, ( + "429 must abort the remaining fan-out" + ) + assert _paths_of(await mp.get_all_model_paths(), "m0") == expected + + +@pytest.mark.asyncio +async def test_openrouter_bad_payload_shapes_preserve_previous_rows( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Malformed successful responses are degraded snapshots, not empty sets.""" + provider = _FakeOpenRouterProvider( + models=[_model("m", canonical_slug="a/m")], db_id=2 + ) + _mock_transport(monkeypatch, lambda request: _endpoints_response("Anthropic")) + await mp.refresh_model_paths([provider]) + before = _paths_of(await mp.get_all_model_paths(), "m") + + for payload in ( + {"data": {"endpoints": None}}, + {"data": {"endpoints": "none"}}, + {"data": {"endpoints": [{"provider_name": "Anthropic"}]}}, + {"data": None}, + {}, + ): + + def _handler( + request: httpx.Request, p: dict[str, Any] | None = payload + ) -> httpx.Response: + return httpx.Response(200, json=p) + + _mock_transport(monkeypatch, _handler) + await mp.refresh_model_paths([provider]) + assert _paths_of(await mp.get_all_model_paths(), "m") == before + + +@pytest.mark.asyncio +async def test_openrouter_shared_base_url_fetched_once( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + """Two providers on the same OpenRouter base URL share the per-cycle + endpoint cache instead of fetching byte-identical bodies twice.""" + native = _FakeOpenRouterProvider( + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=2, + ) + generic = _FakeProvider( + provider_type="generic", + base_url="https://openrouter.ai/api/v1", + models=[_model("claude-opus-4.6", canonical_slug="anthropic/claude-opus-4.6")], + db_id=4, + ) + counter = _mock_transport( + monkeypatch, lambda request: _endpoints_response("Anthropic") + ) + + await mp.refresh_model_paths([native, generic]) + + assert counter["requests"] == 1 + paths = _paths_of(await mp.get_all_model_paths(), "claude-opus-4.6") + assert paths == { + _expected_path(2, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6", "anthropic"), + _expected_path(4, "claude-opus-4.6"), + _expected_path(4, "claude-opus-4.6", "anthropic"), + } + + +@pytest.mark.asyncio +async def test_openrouter_fanout_is_bounded( + patched_session: AsyncEngine, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(mp, "_OPENROUTER_CONCURRENCY", 3) + models = [_model(f"m{i}", canonical_slug=f"author/m{i}") for i in range(20)] + provider = _FakeOpenRouterProvider(models=models, db_id=2) + + state = {"current": 0, "max": 0} + + async def _slow_handler(request: httpx.Request) -> httpx.Response: + state["current"] += 1 + state["max"] = max(state["max"], state["current"]) + await asyncio.sleep(0.02) + state["current"] -= 1 + return _endpoints_response("X") + + def _factory() -> httpx.AsyncClient: + return httpx.AsyncClient(transport=httpx.MockTransport(_slow_handler)) + + monkeypatch.setattr(mp, "_make_http_client", _factory) + + await mp.refresh_model_paths([provider]) + assert state["max"] > 0, "transport fake was never exercised" + assert state["max"] <= 3, f"concurrency exceeded bound: {state['max']}" + + +# --------------------------------------------------------------------------- # +# Query endpoints +# --------------------------------------------------------------------------- # + + +async def _seed_two_provider_shared_model(engine: AsyncEngine) -> None: + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="generic", + base_url="https://other/v1", + models=[_model("claude-opus-4.6")], + db_id=2, + ) + await mp.refresh_model_paths([p1, p2]) + + +@pytest.mark.asyncio +async def test_same_model_two_providers_two_paths( + patched_session: AsyncEngine, +) -> None: + await _seed_two_provider_shared_model(patched_session) + + payload = await mp.get_all_model_paths() + assert len(payload["data"]) == 1 + entry = payload["data"][0] + assert entry["id"] == "claude-opus-4.6" + assert {p["path"] for p in entry["paths"]} == { + _expected_path(1, "claude-opus-4.6"), + _expected_path(2, "claude-opus-4.6"), + } + assert "canonical_id" not in entry + assert all("canonical_id" not in p for p in entry["paths"]) + + +@pytest.mark.asyncio +async def test_get_all_model_paths_keeps_distinct_configured_providers( + patched_session: AsyncEngine, +) -> None: + p1 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("anthropic/claude-opus-4.6")], + db_id=1, + ) + p2 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("claude-opus-4.6")], + db_id=2, + ) + await mp.refresh_model_paths([p1, p2]) + + assert (await mp.get_all_model_paths())["data"] == [ + { + "id": "claude-opus-4.6", + "paths": [ + _path_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ], + } + ] + + +@pytest.mark.asyncio +async def test_get_all_model_paths_is_deterministic( + patched_session: AsyncEngine, +) -> None: + """Output must not depend on rowid insertion order, which changes every + refresh cycle.""" + provider = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("b-model"), _model("a-model")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + first = await mp.get_all_model_paths() + await mp.refresh_model_paths([provider]) + second = await mp.get_all_model_paths() + assert first["data"] == second["data"] + assert [e["id"] for e in first["data"]] == ["a-model", "b-model"] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_returns_route_identity( + patched_session: AsyncEngine, +) -> None: + await _seed_two_provider_shared_model(patched_session) + + payload = await mp.get_paths_for_model("claude-opus-4.6") + assert payload["data"] == [ + _path_entry(1, "claude-opus-4.6"), + _path_entry(2, "claude-opus-4.6"), + ] + assert (await mp.get_paths_for_model("does-not-exist"))["data"] == [] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( + patched_session: AsyncEngine, +) -> None: + provider = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("z-ai/glm-5v-turbo")], + db_id=4, + ) + await mp.refresh_model_paths([provider]) + + assert (await mp.get_paths_for_model("glm-5v-turbo"))["data"] == [ + _path_entry(4, "glm-5v-turbo") + ] + + +@pytest.mark.asyncio +async def test_get_paths_for_model_accepts_provider_prefixed_alias( + patched_session: AsyncEngine, +) -> None: + p1 = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("deepseek-v4-pro")], + db_id=7, + ) + p2 = _FakeProvider( + provider_type="anthropic", + base_url="https://api.anthropic.com/v1", + models=[_model("deepseek/deepseek-v4-pro")], + db_id=4, + ) + await mp.refresh_model_paths([p1, p2]) + + short_paths = (await mp.get_paths_for_model("deepseek-v4-pro"))["data"] + prefixed_paths = (await mp.get_paths_for_model("deepseek/deepseek-v4-pro"))["data"] + + assert short_paths == [ + _path_entry(4, "deepseek-v4-pro"), + _path_entry(7, "deepseek-v4-pro"), + ] + assert prefixed_paths == short_paths + + +@pytest.mark.asyncio +async def test_get_paths_for_model_multi_segment_id_matches_models_listing( + patched_session: AsyncEngine, +) -> None: + """Three-segment upstream IDs resolve to the same first-slash public ID.""" + provider = _FakeProvider( + provider_type="generic", + base_url="https://x/v1", + models=[_model("accounts/fireworks/models/glm-5")], + db_id=1, + ) + await mp.refresh_model_paths([provider]) + + assert _ids_of(await mp.get_all_model_paths()) == {"fireworks/models/glm-5"} + assert (await mp.get_paths_for_model("fireworks/models/glm-5"))["data"] == [ + _path_entry(1, "fireworks/models/glm-5") + ] + assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ + "data" + ] == [_path_entry(1, "fireworks/models/glm-5")] + + +# --------------------------------------------------------------------------- # +# Immediate and periodic refresh +# --------------------------------------------------------------------------- # + + +@pytest.mark.asyncio +async def test_refresh_model_paths_for_provider_selects_mutated_provider( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + target = SimpleNamespace(db_id=2) + other = SimpleNamespace(db_id=1) + seen: list[list[Any]] = [] + + monkeypatch.setattr(proxy, "get_upstreams", lambda: [other, target]) + + async def _fake_refresh(upstreams: list[Any]) -> None: + seen.append(upstreams) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + await mp.refresh_model_paths_for_provider(2) + + assert seen == [[target]] + + +@pytest.mark.asyncio +async def test_admin_refresh_is_disabled_by_model_paths_kill_switch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + calls: list[int] = [] + + async def _fake_refresh(provider_id: int) -> None: + calls.append(provider_id) + + monkeypatch.setattr(mp, "refresh_model_paths_for_provider", _fake_refresh) + await mp.schedule_model_paths_refresh_for_provider(2) + await asyncio.sleep(0) + + assert calls == [] + + +@pytest.mark.asyncio +async def test_admin_refresh_is_backgrounded_and_coalesced( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + mp._scheduled_provider_refresh_ids.clear() + mp._scheduled_provider_refresh_task = None + calls: list[int] = [] + + async def _fake_refresh(provider_id: int) -> None: + calls.append(provider_id) + + monkeypatch.setattr(mp, "refresh_model_paths_for_provider", _fake_refresh) + + await mp.schedule_model_paths_refresh_for_provider(2) + await mp.schedule_model_paths_refresh_for_provider(2) + task = mp._scheduled_provider_refresh_task + assert task is not None + await task + + assert calls == [2] + + +@pytest.mark.asyncio +async def test_refresh_loop_rereads_interval_and_picks_up_providers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 1, raising=False + ) + + seen_batches: list[list[Any]] = [] + + async def _fake_refresh(upstreams: list[Any]) -> None: + seen_batches.append(list(upstreams)) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + sleeps: list[float] = [] + + async def _fast_sleep(seconds: float) -> None: + sleeps.append(seconds) + if len(seen_batches) >= 2: + raise asyncio.CancelledError + + monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) + + batches = [["p1"], ["p1", "p2"]] + + def _provider() -> list[Any]: + return batches[min(len(seen_batches), len(batches) - 1)] + + await mp.refresh_model_paths_periodically(_provider) # type: ignore[arg-type] + + assert seen_batches[0] == ["p1"] + assert seen_batches[1] == ["p1", "p2"], "loop must re-resolve upstreams each cycle" + assert all(s >= 1 for s in sleeps) + + +@pytest.mark.asyncio +async def test_refresh_loop_idles_while_disabled( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A non-positive interval (or the kill switch) must idle the loop, not + exit it, so runtime re-enabling takes effect.""" + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + + refresh_calls: list[Any] = [] + + async def _fake_refresh(upstreams: list[Any]) -> None: + refresh_calls.append(upstreams) + + monkeypatch.setattr(mp, "refresh_model_paths", _fake_refresh) + + idle_sleeps: list[float] = [] + + async def _fast_sleep(seconds: float) -> None: + idle_sleeps.append(seconds) + if len(idle_sleeps) >= 2: + raise asyncio.CancelledError + + monkeypatch.setattr(mp.asyncio, "sleep", _fast_sleep) + + def _upstreams() -> list[BaseUpstreamProvider]: + return [cast(BaseUpstreamProvider, object())] + + await mp.refresh_model_paths_periodically(_upstreams) + + assert refresh_calls == [], "disabled loop must not refresh" + assert len(idle_sleeps) == 2, "disabled loop must keep polling, not exit" + + +def test_refresh_interval_respects_kill_switch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from routstr.core.settings import settings + + monkeypatch.setattr(settings, "enable_model_paths_refresh", False, raising=False) + monkeypatch.setattr( + settings, "model_paths_refresh_interval_seconds", 600, raising=False + ) + assert mp._refresh_interval_seconds() == 0 + + monkeypatch.setattr(settings, "enable_model_paths_refresh", True, raising=False) + assert mp._refresh_interval_seconds() == 600 + + +# --------------------------------------------------------------------------- # +# HTTP endpoints +# --------------------------------------------------------------------------- # + + +def _make_model_paths_app() -> FastAPI: + app = FastAPI() + app.include_router(models_router) + return app + + +def test_model_paths_endpoint_returns_all_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + expected = { + "data": [ + { + "id": "claude-opus-4.6", + "paths": [ + _path_entry(1, "claude-opus-4.6"), + _path_entry( + 2, + "claude-opus-4.6", + endpoint_tag="google-vertex/us", + endpoint_name="Google", + ), + ], + } + ], + "updated_at": 1753500000, + } + + async def _fake_get_all_model_paths() -> dict[str, Any]: + return expected + + monkeypatch.setattr(mp, "get_all_model_paths", _fake_get_all_model_paths) + + response = TestClient(_make_model_paths_app()).get("/v1/models/paths") + + assert response.status_code == 200 + assert response.json() == expected + + +def test_model_paths_for_model_returns_404_for_unknown_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: []) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "does-not-exist"} + ) + + assert response.status_code == 404 + assert response.json() == {"detail": "Model not found"} + + +def test_model_paths_for_known_model_can_return_empty_collection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: [_model("known")]) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "known"} + ) + + assert response.status_code == 200 + assert response.json() == {"data": [], "updated_at": None} + + +def test_model_paths_for_routing_only_alias_returns_404( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import routstr.proxy as proxy + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + return {"data": [], "updated_at": None} + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + monkeypatch.setattr(proxy, "get_unique_models", lambda: [_model("advertised")]) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", params={"model_id": "routing-alias"} + ) + + assert response.status_code == 404 + + +def test_model_paths_for_model_endpoint_accepts_slash_model_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[str] = [] + + expected = { + "data": [ + _path_entry( + 2, + "anthropic/claude-opus-4.6", + endpoint_tag="anthropic", + endpoint_name="Anthropic", + ) + ], + "updated_at": None, + } + + async def _fake_get_paths_for_model(model_id: str) -> dict[str, Any]: + calls.append(model_id) + return expected + + monkeypatch.setattr(mp, "get_paths_for_model", _fake_get_paths_for_model) + + response = TestClient(_make_model_paths_app()).get( + "/v1/models/paths/model", + params={"model_id": "anthropic/claude-opus-4.6"}, + ) + + assert response.status_code == 200 + assert response.json() == expected + assert calls == ["anthropic/claude-opus-4.6"] diff --git a/tests/unit/test_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_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_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 diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 99abc17d..31dd5767 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -71,7 +71,9 @@ async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None: @pytest.mark.asyncio -async def test_pay_for_request_sets_reserved_at_on_child_key(session: AsyncSession) -> None: +async def test_pay_for_request_sets_reserved_at_on_child_key( + session: AsyncSession, +) -> None: parent = ApiKey(hashed_key="parentkey", balance=10_000) child = ApiKey(hashed_key="childkey", balance=0, parent_key_hash="parentkey") session.add(parent) @@ -204,7 +206,9 @@ async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) -> @pytest.mark.asyncio -async def test_release_stale_reservations_skips_null_reserved_at(session: AsyncSession) -> None: +async def test_release_stale_reservations_skips_null_reserved_at( + session: AsyncSession, +) -> None: # Reservations without a timestamp may belong to instances running older # code (rolling deploy) — the background sweeper must not touch them. key = ApiKey( @@ -224,7 +228,9 @@ async def test_release_stale_reservations_skips_null_reserved_at(session: AsyncS @pytest.mark.asyncio -async def test_reset_all_reserved_balances_clears_reserved_at(session: AsyncSession) -> None: +async def test_reset_all_reserved_balances_clears_reserved_at( + session: AsyncSession, +) -> None: key = ApiKey( hashed_key="resetkey", balance=5_000, @@ -404,9 +410,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None: AsyncMock(return_value=1_000), ), patch.object(proxy_module, "check_token_balance", MagicMock()), - patch.object( - proxy_module, "get_bearer_token_key", AsyncMock(return_value=key) - ), + patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)), patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)), patch.object( proxy_module, @@ -418,6 +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, 1_000, 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 216e1c3a..495f1e57 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -383,9 +383,7 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None: AsyncMock(return_value=1_000), ), patch.object(proxy_module, "check_token_balance", MagicMock()), - patch.object( - proxy_module, "get_bearer_token_key", AsyncMock(return_value=key) - ), + patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)), patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)), patch.object( proxy_module, @@ -408,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, 1_000, 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 679b7d65..1a8c1f82 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -2,10 +2,13 @@ import asyncio import base64 import json import socket +from collections.abc import AsyncIterator, Generator +from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest +from cashu.core.base import MeltQuoteState from routstr.core.db import ApiKey from routstr.wallet import ( @@ -13,6 +16,7 @@ from routstr.wallet import ( Bolt11PaymentNotAttempted, MintConnectionError, TokenConsumedError, + _is_mint_rate_limited, classify_redemption_error, credit_balance, execute_bolt11_payment, @@ -25,6 +29,26 @@ from routstr.wallet import ( ) +@pytest.fixture(autouse=True) +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 + + 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_concurrency = original_concurrency + wallet_module._MintRateGuard._guards.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() @@ -34,13 +58,57 @@ async def test_get_balance() -> None: # Reset the module-level wallet cache so a real wallet cached by an earlier # test (e.g. an unmocked admin-withdraw path) can't shadow the mock here. - with patch("routstr.wallet._wallets", {}), patch( - "routstr.wallet.Wallet.with_db", return_value=mock_wallet + with ( + patch("routstr.wallet._wallets", {}), + patch("routstr.wallet.Wallet.with_db", return_value=mock_wallet), ): balance = await get_balance("sat") 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 = { @@ -147,6 +215,166 @@ 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 + from routstr.wallet import SourceMintConnectionError + + 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("mint unavailable")) + ) + get_wallet = AsyncMock(return_value=source_wallet) + + with ( + patch.object(settings, "primary_mint", source), + patch.object(settings, "cashu_mints", [source, destination]), + patch("routstr.wallet.deserialize_token_from_string", return_value=token), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.logger.warning") as warning, + ): + with pytest.raises(SourceMintConnectionError): + await recieve_token("cashuAtoken") + + get_wallet.assert_awaited_once_with(source, "sat", load=False) + failure = next( + call.kwargs["extra"] + for call in warning.call_args_list + 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 +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() @@ -157,10 +385,128 @@ 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], + load_proofs=AsyncMock(), + 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.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 + + +@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_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 + preferred = "http://preferred:3338" + primary = "http://primary:3338" preferred_wallet = Mock(keysets={}, proofs=[]) preferred_wallet.select_to_send = AsyncMock() primary_wallet = Mock(keysets={}, proofs=[]) @@ -173,24 +519,37 @@ async def test_send_falls_back_when_preferred_mint_has_only_reserved_balance() - primary_liquid = Mock(amount=1000, reserved=False) primary_wallet.select_to_send.return_value = ([primary_liquid], None) - async def get_wallet(mint_url: str, unit: str) -> Mock: + async def get_wallet(mint_url: str, unit: str, **_: object) -> Mock: assert unit == "sat" - return primary_wallet if mint_url == "http://primary:3338" else preferred_wallet + return primary_wallet if mint_url == primary else preferred_wallet - def get_proofs(wallet: Mock, mint_url: str, unit: str) -> list[Mock]: + def get_proofs( + wallet: Mock, + mint_url: str, + unit: str, + *, + not_reserved: bool = False, + ) -> list[Mock]: assert unit == "sat" if wallet is primary_wallet: - assert mint_url == "http://primary:3338" - return [primary_liquid] - assert mint_url == "http://preferred:3338" - return [preferred_liquid, preferred_reserved] + assert mint_url == primary + proofs = [primary_liquid] + else: + assert mint_url == preferred + proofs = [preferred_liquid, preferred_reserved] + return ( + [proof for proof in proofs if not proof.reserved] + if not_reserved + else proofs + ) with ( - patch.object(settings, "primary_mint", "http://primary:3338"), + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary, preferred]), patch("routstr.wallet.get_wallet", side_effect=get_wallet), patch("routstr.wallet.get_proofs_per_mint_and_unit", side_effect=get_proofs), ): - amount, token = await send(1000, "sat", "http://preferred:3338") + amount, token = await send(1000, "sat", preferred) assert (amount, token) == (1000, "primary-token") preferred_wallet.select_to_send.assert_not_awaited() @@ -203,21 +562,30 @@ async def test_send_falls_back_when_preferred_mint_has_only_reserved_balance() - async def test_send_primary_with_only_reserved_proofs_still_raises() -> None: from routstr.core.settings import settings + primary = "http://primary:3338" wallet = Mock(keysets={}, proofs=[]) - wallet.select_to_send = AsyncMock(side_effect=RuntimeError("balance too low")) + wallet.select_to_send = AsyncMock() reserved = Mock(amount=1000, reserved=True) - with ( - patch.object(settings, "primary_mint", "http://primary:3338"), - patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), - patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[reserved]), - pytest.raises(RuntimeError, match="balance too low"), - ): - await send(1000, "sat", "http://primary:3338") + def get_proofs( + _wallet: Mock, + _mint_url: str, + _unit: str, + *, + not_reserved: bool = False, + ) -> list[Mock]: + return [] if not_reserved else [reserved] - wallet.select_to_send.assert_awaited_once_with( - [], 1000, set_reserved=False, include_fees=False - ) + with ( + patch.object(settings, "primary_mint", primary), + patch.object(settings, "cashu_mints", [primary]), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.wallet.get_proofs_per_mint_and_unit", side_effect=get_proofs), + pytest.raises(ValueError, match="No trusted mint has"), + ): + await send(1000, "sat", primary) + + wallet.select_to_send.assert_not_awaited() @pytest.mark.asyncio @@ -258,6 +626,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 @@ -391,7 +780,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") @@ -503,7 +892,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 @@ -634,7 +1025,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 @@ -665,7 +1056,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 @@ -801,9 +1192,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 @@ -828,9 +1217,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 @@ -878,12 +1265,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 @@ -1195,6 +1578,21 @@ 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_rate_limited", + 503, + "Cashu mint rate-limited; retry after cooldown", + "cashu_mint_rate_limited", + ) + + @pytest.mark.parametrize( "error", [ @@ -1231,7 +1629,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"), ], @@ -1286,7 +1686,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) == ( @@ -1349,7 +1751,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 @@ -1373,16 +1777,24 @@ 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") + 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 @@ -1390,7 +1802,7 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None: 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 @@ -1608,3 +2020,881 @@ async def test_execute_bolt11_payment_rereserves_when_cancelled() -> None: plan.wallet.set_reserved_for_melt.assert_awaited_once_with( plan.proofs, reserved=True, quote_id="quote-1" ) + + +# --------------------------------------------------------------------------- +# Per-mint adaptive guard + _mint_operation factory/retry +# --------------------------------------------------------------------------- + + +@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_guard_bounds_concurrency() -> None: + from routstr.wallet import _MintRateGuard + + guard = _MintRateGuard("http://mint:3338", 2) + active = 0 + peak = 0 + + async def operation() -> None: + nonlocal active, peak + active += 1 + peak = max(peak, active) + await asyncio.sleep(0) + active -= 1 + + await asyncio.gather(*(guard.run(operation) for _ in range(5))) + + assert peak == 2 + + +@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.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) + 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.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 + 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 + + 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 + + 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_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.mint.time.monotonic", return_value=0), + patch("routstr.mint.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 + 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.object(settings, "mint_max_concurrency", 1): + 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" + ) + + assert result == "ok" + sleep.assert_awaited_once_with(60.0) + + +@pytest.mark.asyncio +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="ok") + with ( + patch.object(settings, "mint_max_concurrency", 1), + patch.object(settings, "mint_operation_timeout_seconds", 0.01), + patch("routstr.mint.asyncio.sleep", AsyncMock()) as sleep, + ): + guard = _MintRateGuard.get("http://mint:3338") + guard.apply_cooldown(60) + assert await _mint_operation(operation, mint_url="http://mint:3338") == "ok" + + 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.mint.asyncio.sleep", AsyncMock()), + ): + assert await _mint_operation(operation, mint_url="http://mint:3338") == "ok" + + assert operation.await_count == 2 + + +@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.mint.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: + # A fresh wallet must load even when the host has been up for less than + # the reload interval. + 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") + ) + + 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_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.mint.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. + 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() -> str: + 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_concurrency", 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() -> None: + 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_concurrency", 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_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( + 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_when_primary_wallet_cannot_load() -> None: + from routstr.core.settings import settings + from routstr.wallet import swap_to_primary_mint + + primary = "http://primary:3338" + secondary = "http://secondary:3338" + foreign = "http://foreign:3338" + + token = Mock( + mint=foreign, + unit="sat", + amount=1000, + keysets=["keyset1"], + proofs=[Mock(amount=1000)], + ) + 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(state=MeltQuoteState.paid)), + ) + + mint_quote = Mock(quote="mint_q_secondary", request="lnbc1secondary") + 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()), + ) + + async def get_wallet(mint: str, *args: object, **kwargs: object) -> Mock: + if mint == primary: + raise httpx.ConnectError("primary down") + return secondary_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, + patch("routstr.wallet.logger.info") as info, + ): + amount, unit, mint_url = await swap_to_primary_mint(token, source_wallet) + + 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, *info.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 +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" + + 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=captured.value) + + 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_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(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_concurrency", 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) + + +@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.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]), + ), + ): + _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() + + +# --------------------------------------------------------------------------- +# _is_mint_rate_limited — strict HTTP 429 only (no substring matching) +# --------------------------------------------------------------------------- + + +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 backoff escalation and recovery +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_probe_escalates_consecutive_rate_limits() -> None: + from routstr.mint import MintRateGuard + + guard = MintRateGuard("http://mint", max_concurrency=0) + guard.apply_rate_limit_cooldown() + guard._cooldown_until = 0.0 + + with pytest.raises(httpx.HTTPStatusError): + await guard.run(AsyncMock(side_effect=_http_429_error())) + + assert guard._consecutive_rate_limits == 2 + assert guard._needs_probe is True + assert guard.cooldown_remaining() > 60 + + +@pytest.mark.asyncio +async def test_probe_recovery_resets_consecutive_rate_limits() -> None: + from routstr.mint import MintRateGuard + + guard = MintRateGuard("http://mint", max_concurrency=0) + guard.apply_rate_limit_cooldown() + guard._cooldown_until = 0.0 + + assert await guard.run(AsyncMock(return_value="ok")) == "ok" + + 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 + ) 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 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..b97e0e72 100644 --- a/ui/lib/api/services/wallet.ts +++ b/ui/lib/api/services/wallet.ts @@ -36,10 +36,13 @@ 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 { token: string; + mint_url: string; } export interface CreateChildKeyResponse {