diff --git a/.env.example b/.env.example index 8ff04b35..9f0c0bbc 100644 --- a/.env.example +++ b/.env.example @@ -22,6 +22,19 @@ ROUTSTR_SECRET_KEY= # Database # DATABASE_URL=sqlite+aiosqlite:///keys.db +# Pool controls are validated at boot, sourced only from the environment, and +# 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_RECYCLE=1800 +# DATABASE_POOL_PRE_PING=false +# Warn when a checkout is held this many seconds. +# DATABASE_POOL_HOLD_WARN_SECONDS=10 +# SQLite serialises writes; increasing its pool can trade pool timeouts for +# "database is locked" errors rather than increasing write throughput. # Node Information # NAME=My Routstr Node @@ -31,7 +44,9 @@ ROUTSTR_SECRET_KEY= # RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com" # ENABLE_ANALYTICS_SHARING=true # CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me" +# MINT_OPERATION_CONCURRENCY=4 # RECEIVE_LN_ADDRESS= +# REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS=900 # Custom Pricing Configuration # MODEL_BASED_PRICING=true diff --git a/migrations/versions/aa50fde387a2_add_refund_sweep_claim_lease.py b/migrations/versions/aa50fde387a2_add_refund_sweep_claim_lease.py new file mode 100644 index 00000000..40c7e577 --- /dev/null +++ b/migrations/versions/aa50fde387a2_add_refund_sweep_claim_lease.py @@ -0,0 +1,39 @@ +"""add refund sweep claim lease + +Revision ID: aa50fde387a2 +Revises: 9c4d8e2f1a6b +Create Date: 2026-07-26 12:50:10.509217 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +revision = "aa50fde387a2" +down_revision = "9c4d8e2f1a6b" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + columns = { + column["name"] + for column in sa.inspect(conn).get_columns("cashu_transactions") + } + if "sweep_started_at" not in columns: + op.add_column( + "cashu_transactions", + sa.Column("sweep_started_at", sa.Integer(), nullable=True), + ) + + +def downgrade() -> None: + conn = op.get_bind() + columns = { + column["name"] + for column in sa.inspect(conn).get_columns("cashu_transactions") + } + if "sweep_started_at" in columns: + op.drop_column("cashu_transactions", "sweep_started_at") diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 66a1d288..74a94668 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1669,12 +1669,18 @@ async def get_transactions_api( ) total = count_result.one() - stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit) + stmt = ( + base.order_by(col(CashuTransaction.created_at).desc()) + .offset(offset) + .limit(limit) + ) results = await session.exec(stmt) transactions = results.all() return { - "transactions": [tx.dict() for tx in transactions], + "transactions": [ + tx.dict(exclude={"sweep_started_at"}) for tx in transactions + ], "total": total, } diff --git a/routstr/core/db.py b/routstr/core/db.py index 16cfedfd..12dcc659 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -12,21 +12,80 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from alembic.util.exc import CommandError -from sqlalchemy import Index, UniqueConstraint, case, delete, or_ +from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_ +from sqlalchemy.engine import make_url from sqlalchemy.exc import IntegrityError, OperationalError +from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlalchemy.orm import aliased from sqlmodel import Field, Relationship, SQLModel, col, func, select, update from sqlmodel.ext.asyncio.session import AsyncSession from .logging import get_logger +from .settings import settings logger = get_logger(__name__) DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") -engine = create_async_engine(DATABASE_URL, echo=False) # echo=True for debugging SQL +def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: + """Build and instrument an async engine from environment-only settings.""" + url = make_url(database_url) + backend = url.get_backend_name() + is_sqlite = backend == "sqlite" + is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"} + pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite + options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping} + if not is_memory_sqlite: + options.update( + pool_size=settings.database_pool_size, + max_overflow=settings.database_max_overflow, + pool_timeout=settings.database_pool_timeout, + pool_recycle=settings.database_pool_recycle, + ) + + logger.info( + "Database pool configured", + extra={ + "database_url_backend": backend, + "in_memory_sqlite": is_memory_sqlite, + **options, + }, + ) + created_engine = create_async_engine(database_url, echo=False, **options) + hold_warn_seconds = settings.database_pool_hold_warn_seconds + + def record_pool_checkout( + dbapi_connection: object, connection_record: object, proxy: object + ) -> None: + connection_record.info["routstr_checked_out_at"] = time.monotonic() # type: ignore[attr-defined] + + def record_pool_checkin( + dbapi_connection: object, connection_record: object + ) -> None: + checked_out_at = connection_record.info.pop( # type: ignore[attr-defined] + "routstr_checked_out_at", None + ) + if checked_out_at is None: + return + held_seconds = time.monotonic() - checked_out_at + if held_seconds >= hold_warn_seconds: + logger.warning( + "Database connection held longer than threshold", + extra={ + "held_seconds": round(held_seconds, 3), + "threshold_seconds": hold_warn_seconds, + "pool_status": created_engine.pool.status(), + }, + ) + + event.listen(created_engine.sync_engine, "checkout", record_pool_checkout) + event.listen(created_engine.sync_engine, "checkin", record_pool_checkin) + return created_engine + + +engine = create_db_engine() class ApiKey(SQLModel, table=True): # type: ignore @@ -222,7 +281,10 @@ async def release_stale_reservations( if released: logger.warning( "Released stale reservations", - extra={"released_reservations": released, "max_age_seconds": max_age_seconds}, + extra={ + "released_reservations": released, + "max_age_seconds": max_age_seconds, + }, ) return released @@ -320,7 +382,8 @@ class LightningInvoice(SQLModel, table=True): # type: ignore description: str = Field(description="Invoice description") payment_hash: str = Field(description="Payment hash for tracking", unique=True) status: str = Field( - default="pending", description="pending, paid, expired, cancelled" + default="pending", + description="pending, paid, expired, cancelled, reconciliation_required", ) api_key_hash: str | None = Field( default=None, description="Associated API key hash for topup operations" @@ -365,6 +428,10 @@ class CashuTransaction(SQLModel, table=True): # type: ignore ) collected: bool = Field(default=False) swept: bool = Field(default=False) + sweep_started_at: int | None = Field( + default=None, + description="Unix timestamp for a recoverable refund-sweep claim", + ) source: str = Field( default="x-cashu", description="Payment source: x-cashu or apikey", @@ -717,14 +784,43 @@ async def complete_routstr_fee_payout( return result.rowcount == 1 -async def balances_for_mint_and_unit( +async def balance_for_mint_and_unit( db_session: AsyncSession, mint_url: str, unit: str ) -> int: - query = select(func.sum(ApiKey.balance)).where( - ApiKey.refund_mint_url == mint_url, ApiKey.refund_currency == unit + """Return the user liability for one mint and unit in millisatoshis.""" + result = await db_session.exec( + select(func.sum(ApiKey.balance)).where( + col(ApiKey.refund_mint_url) == mint_url, + col(ApiKey.refund_currency) == unit, + ) + ) + return int(result.one() or 0) + + +async def balances_by_mint_and_unit( + db_session: AsyncSession, mint_urls: list[str], units: list[str] +) -> dict[tuple[str, str], int]: + """Return requested user liabilities grouped by mint and unit.""" + if not mint_urls or not units: + return {} + query = ( + select( + col(ApiKey.refund_mint_url), + col(ApiKey.refund_currency), + func.sum(ApiKey.balance), + ) + .where( + col(ApiKey.refund_mint_url).in_(mint_urls), + col(ApiKey.refund_currency).in_(units), + ) + .group_by(col(ApiKey.refund_mint_url), col(ApiKey.refund_currency)) ) result = await db_session.exec(query) - return result.one() or 0 + return { + (mint_url, unit): int(balance or 0) + for mint_url, unit, balance in result.all() + if mint_url is not None and unit is not None + } async def init_db() -> None: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index a0b5d05c..0caffd16 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -41,6 +41,9 @@ class Settings(BaseSettings): receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS") primary_mint: str = Field(default="", env="PRIMARY_MINT_URL") primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT") + mint_operation_concurrency: int = Field( + default=4, ge=1, env="MINT_OPERATION_CONCURRENCY" + ) # Lightning payout configuration # Minimum available balance (in satoshis) before profit is paid out over @@ -99,6 +102,24 @@ class Settings(BaseSettings): enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS") + refund_sweep_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_pool_timeout: float = Field( + default=30.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") + database_pool_hold_warn_seconds: float = Field( + default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS" + ) # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") @@ -144,10 +165,32 @@ def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]: # ``routstr.core.vault``. SECRET_FIELDS = frozenset({"admin_password", "nsec"}) +# Infrastructure the node needs *before* it can open a DB session — so it can +# never be configured from the DB (chicken-and-egg) and stays env-only. Unlike +# secrets (owned by bootstrap), these are excluded so the DB settings blob can +# neither store nor shadow them; env is always authoritative. +ENV_ONLY_FIELDS = frozenset( + { + "database_pool_size", + "database_max_overflow", + "database_pool_timeout", + "database_pool_recycle", + "database_pool_pre_ping", + "database_pool_hold_warn_seconds", + } +) + +_NON_PERSISTED_FIELDS = SECRET_FIELDS | ENV_ONLY_FIELDS + def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]: - """Return a copy of ``data`` without any secret fields (for persistence).""" - return {k: v for k, v in data.items() if k not in SECRET_FIELDS} + """Return a copy of ``data`` without secret or env-only fields. + + Both are kept out of the persisted settings blob: secrets for confidentiality, + env-only fields (e.g. DB pool sizing) because they must never be sourced from + the database. + """ + return {k: v for k, v in data.items() if k not in _NON_PERSISTED_FIELDS} def _apply_to_live_settings(data: dict[str, Any]) -> None: @@ -327,7 +370,13 @@ class SettingsService: valid_fields = set(env_resolved.dict().keys()) merged_dict: dict[str, Any] = dict(env_resolved.dict()) merged_dict.update( - {k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields} + { + k: v + for k, v in db_json.items() + if v not in (None, "", [], {}) + and k in valid_fields + and k not in ENV_ONLY_FIELDS + } ) merged_dict = Settings(**merged_dict).dict() @@ -394,8 +443,13 @@ class SettingsService: ) ) await db_session.commit() - # Update in-place + # Update in-place. Env-only fields (e.g. DB pool sizing) are never + # applied here: the engine pool is already built at boot from env, + # so letting an update mutate the live value would only make it + # diverge from the running pool. for k, v in candidate.dict().items(): + if k in ENV_ONLY_FIELDS: + continue setattr(settings, k, v) cls._current = settings return settings diff --git a/routstr/lightning.py b/routstr/lightning.py index b0bbc63b..d696403a 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -5,7 +5,8 @@ import time from fastapi import APIRouter, Depends, Header, HTTPException from pydantic import BaseModel, Field -from sqlmodel import col, select +from sqlalchemy.orm.attributes import set_committed_value +from sqlmodel import col, select, update from sqlmodel.ext.asyncio.session import AsyncSession from .core.db import ApiKey, LightningInvoice, create_session, get_session @@ -222,80 +223,170 @@ async def recover_invoice( async def check_invoice_payment( 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() + wallet = await get_wallet(settings.primary_mint, "sat") + mint_status = await wallet.get_mint_quote(invoice_payment_hash) + if not mint_status.paid: + return - mint_status = await wallet.get_mint_quote(invoice.payment_hash) + # 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") + ) + 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 + ) + logger.critical( + "Paid topup invoice target API key was not found; reconciliation required", + extra={"invoice_id": invoice_id}, + ) + return - if mint_status.paid: - invoice.status = "paid" - invoice.paid_at = int(time.time()) + # 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 - if invoice.purpose == "create": - api_key = await create_api_key_from_invoice(invoice, session) - invoice.api_key_hash = api_key.hashed_key - elif invoice.purpose == "topup" and invoice.api_key_hash: - await topup_api_key_from_invoice(invoice, session) + # 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) - await session.commit() - - logger.info( - "Lightning invoice paid", - extra={ - "invoice_id": invoice.id, - "amount_sats": invoice.amount_sats, - "purpose": invoice.purpose, - "api_key_hash": invoice.api_key_hash[:8] + "..." - if invoice.api_key_hash - else None, - }, + # 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, + ) ) + 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 + ) + set_committed_value(invoice, "status", committed_invoice.status) + set_committed_value(invoice, "paid_at", committed_invoice.paid_at) + return + await finalization_session.commit() - except Exception as e: + # 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, + }, + ) + 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): + raise logger.error(f"Failed to check invoice payment: {e}") -async def create_api_key_from_invoice( +async def _create_api_key_record( invoice: LightningInvoice, session: AsyncSession ) -> ApiKey: - wallet = await get_wallet(settings.primary_mint, "sat") - await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) - dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}" hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest() - api_key = ApiKey( hashed_key=hashed_key, - balance=invoice.amount_sats * 1000, # Convert to msats + balance=invoice.amount_sats * 1000, refund_currency="sat", refund_mint_url=settings.primary_mint, balance_limit=invoice.balance_limit, balance_limit_reset=invoice.balance_limit_reset, validity_date=invoice.validity_date, ) - session.add(api_key) await session.flush() - return api_key -async def topup_api_key_from_invoice( +async def _credit_topup_record( invoice: LightningInvoice, session: AsyncSession ) -> None: - wallet = await get_wallet(settings.primary_mint, "sat") - await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) - if not invoice.api_key_hash: raise ValueError("No API key associated with topup invoice") - - api_key = await session.get(ApiKey, invoice.api_key_hash) - if not api_key: + credited = 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) + ) + if credited.rowcount != 1: raise ValueError("Associated API key not found") - api_key.balance += invoice.amount_sats * 1000 # Convert to msats - await session.flush() - INVOICE_WATCH_INTERVAL_SECONDS = 5 INVOICE_WATCH_BATCH_LIMIT = 100 diff --git a/routstr/proxy.py b/routstr/proxy.py index c0b8794f..46596077 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -1,4 +1,5 @@ import asyncio +import inspect import json from typing import Any @@ -220,6 +221,20 @@ _API_PATH_PREFIXES = ( @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) +) -> Response | StreamingResponse: + """Run proxy setup in a short request session, never across response streaming.""" + try: + return await _proxy(request, path, session) + finally: + # FastAPI yield dependencies normally close after the response body is + # sent. Close explicitly so a long stream cannot retain DB resources. + close_result = session.close() + if inspect.isawaitable(close_result): + await close_result + + +async def _proxy( + request: Request, path: str, session: AsyncSession ) -> Response | StreamingResponse: # GET requests must hit a known API prefix; otherwise return a 404 (HTML # for browsers, JSON for API clients). POST requests are always forwarded diff --git a/routstr/wallet.py b/routstr/wallet.py index dd92d913..aed99f92 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -33,6 +33,22 @@ _CashuMintInfo.model_rebuild(force=True) logger = get_logger(__name__) +def _sats_to_msats(amount: int) -> int: + return amount * 1000 + + +def _msats_to_sats(amount: int) -> int: + return amount // 1000 + + +def _mints_to_inspect() -> list[str]: + """Return configured mints plus the primary mint, without duplicates.""" + mint_urls = list(settings.cashu_mints) + if settings.primary_mint and settings.primary_mint not in mint_urls: + mint_urls.append(settings.primary_mint) + return mint_urls + + class MintConnectionError(Exception): """The mint could not be reached (network transport failure). @@ -305,10 +321,10 @@ def _net_minted_amount(amount_msat: int, token_unit: str, fees: int) -> int: Convert the token value minus fees (given in the token unit) into an amount in the primary mint's unit. """ - fee_msat = fees * 1000 if token_unit == "sat" else fees + fee_msat = _sats_to_msats(fees) if token_unit == "sat" else fees remaining_msat = amount_msat - fee_msat if settings.primary_mint_unit == "sat": - return int(remaining_msat // 1000) + return _msats_to_sats(remaining_msat) return int(remaining_msat) @@ -361,7 +377,7 @@ async def _calculate_swap_amount( melt fees and NUT-02 input fees on the foreign mint. """ if settings.primary_mint_unit == "sat": - receive_amount = amount_msat // 1000 + receive_amount = _msats_to_sats(amount_msat) else: receive_amount = amount_msat @@ -395,7 +411,7 @@ async def _calculate_swap_amount( logger.info( "swap_to_primary_mint: fee estimation result", extra={ - "token_amount_sat": amount_msat // 1000, + "token_amount_sat": _msats_to_sats(amount_msat), "estimated_fee": total_fees, "estimated_fee_unit": token_unit, "input_fees": input_fees, @@ -434,7 +450,7 @@ async def swap_to_primary_mint( token_amount = token_obj.amount if token_obj.unit == "sat": - amount_msat = token_amount * 1000 + amount_msat = _sats_to_msats(token_amount) elif token_obj.unit == "msat": amount_msat = token_amount else: @@ -598,7 +614,10 @@ async def swap_to_primary_mint( # advance the counter so the next request derives fresh secrets. logger.warning( "swap_to_primary_mint: outputs already signed — recovering orphaned proofs", - extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount}, + extra={ + "mint_quote_id": mint_quote.quote, + "minted_amount": minted_amount, + }, ) try: for keyset_id in primary_wallet.keysets: @@ -683,7 +702,7 @@ async def credit_balance( ) if unit == "sat": - amount = amount * 1000 + amount = _sats_to_msats(amount) logger.info( "credit_balance: Converted to msat", extra={"amount_msat": amount} ) @@ -836,18 +855,34 @@ async def fetch_all_balances( if units is None: units = ["sat", "msat"] - async def fetch_balance( - session: db.AsyncSession, mint_url: str, unit: str - ) -> BalanceDetail: - try: - wallet = await get_wallet(mint_url, unit) - proofs = get_proofs_per_mint_and_unit( - wallet, mint_url, unit, not_reserved=True + # 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. + mint_urls = _mints_to_inspect() + + user_balances: dict[tuple[str, str], int] = {} + liabilities_error: str | None = None + try: + async with db.create_session() as session: + user_balances = await db.balances_by_mint_and_unit( + session, mint_urls, units ) - proofs = await slow_filter_spend_proofs(proofs, wallet) - user_balance = await db.balances_for_mint_and_unit(session, mint_url, unit) + except Exception as e: + logger.error("Error reading user balances", extra={"error": str(e)}) + liabilities_error = str(e) + + 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 + ) + proofs = await slow_filter_spend_proofs(proofs, wallet) + user_balance = user_balances.get((mint_url, unit), 0) if unit == "sat": - user_balance = user_balance // 1000 + user_balance = _msats_to_sats(user_balance) proofs_balance = sum(proof.amount for proof in proofs) result: BalanceDetail = { @@ -870,48 +905,37 @@ async def fetch_all_balances( } return error_result - # Build the set of mints to inspect. Received tokens are stored against - # ``primary_mint`` (which defaults to a real mint even when ``cashu_mints`` - # is empty), so include it as a fallback — otherwise a node that accepts - # payments would still report empty balances when ``cashu_mints`` is unset. - mint_urls: list[str] = list(settings.cashu_mints) - if settings.primary_mint and settings.primary_mint not in mint_urls: - mint_urls.append(settings.primary_mint) + tasks = [fetch_balance(mint_url, unit) for mint_url in mint_urls for unit in units] + balance_details = list(await asyncio.gather(*tasks)) - # Create tasks for all mint/unit combinations - async with db.create_session() as session: - tasks = [ - fetch_balance(session, mint_url, unit) - for mint_url in mint_urls - for unit in units - ] - - # Run all tasks concurrently - balance_details = list(await asyncio.gather(*tasks)) - - # Calculate totals total_wallet_balance_sats = 0 total_user_balance_sats = 0 - for detail in balance_details: - if not detail.get("error"): - # Convert to sats for total calculation - unit = detail["unit"] - proofs_balance_sats = ( - detail["wallet_balance"] - if unit == "sat" - else detail["wallet_balance"] // 1000 - ) - user_balance_sats = ( + if detail.get("error"): + continue + unit = detail["unit"] + total_wallet_balance_sats += ( + detail["wallet_balance"] + if unit == "sat" + else _msats_to_sats(detail["wallet_balance"]) + ) + if liabilities_error is None: + total_user_balance_sats += ( detail["user_balance"] if unit == "sat" - else detail["user_balance"] // 1000 + else _msats_to_sats(detail["user_balance"]) ) - total_wallet_balance_sats += proofs_balance_sats - total_user_balance_sats += user_balance_sats - - owner_balance = total_wallet_balance_sats - total_user_balance_sats + 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 + detail["owner_balance"] = 0 + detail.setdefault("error", liabilities_error) return ( balance_details, @@ -931,61 +955,87 @@ async def periodic_payout() -> None: # Include the primary mint even if it is not listed in cashu_mints, # matching fetch_all_balances(); otherwise primary-mint funds never # auto-payout. - mint_urls: list[str] = list(settings.cashu_mints) - if settings.primary_mint and settings.primary_mint not in mint_urls: - mint_urls.append(settings.primary_mint) + mint_urls = _mints_to_inspect() - async with db.create_session() as session: - for mint_url in mint_urls: - for unit in ["sat", "msat"]: - # Isolate failures per mint/unit so one slow or failing - # mint does not abort payout for every other mint/unit. - try: - wallet = await get_wallet(mint_url, unit) - proofs = get_proofs_per_mint_and_unit( - wallet, mint_url, unit, not_reserved=True - ) - proofs = await slow_filter_spend_proofs(proofs, wallet) - await asyncio.sleep(5) - user_balance = await db.balances_for_mint_and_unit( + for mint_url in mint_urls: + for unit in ["sat", "msat"]: + # Isolate failures per mint/unit so one slow or failing + # mint does not abort payout for every other mint/unit. + try: + wallet = await get_wallet(mint_url, unit) + proofs = get_proofs_per_mint_and_unit( + wallet, mint_url, unit, not_reserved=True + ) + proofs = await slow_filter_spend_proofs(proofs, wallet) + await asyncio.sleep(5) + except Exception as e: + logger.error( + f"Error sending payout: {type(e).__name__}", + extra={ + "error": str(e), + "mint_url": mint_url, + "unit": unit, + }, + ) + continue + + # Fetch the liability AFTER the proofs snapshot and the + # settle delay: a concurrent top-up then only inflates the + # liability, shrinking the payout — never sending + # customer-backed funds as profit. + try: + async with db.create_session() as session: + user_balance = await db.balance_for_mint_and_unit( session, mint_url, unit ) - if unit == "sat": - user_balance = user_balance // 1000 - proofs_balance = sum(proof.amount for proof in proofs) - available_balance = proofs_balance - user_balance - # Threshold is configured in sats; convert for msat wallets. - min_amount = ( - settings.min_payout_sat - if unit == "sat" - else settings.min_payout_sat * 1000 + except Exception as e: + logger.error( + f"Error in periodic payout cycle: {type(e).__name__}", + extra={ + "error": str(e), + "mint_url": mint_url, + "unit": unit, + }, + ) + continue + + try: + if unit == "sat": + user_balance = _msats_to_sats(user_balance) + proofs_balance = sum(proof.amount for proof in proofs) + available_balance = proofs_balance - user_balance + # Threshold is configured in sats; convert for msat wallets. + min_amount = ( + settings.min_payout_sat + if unit == "sat" + else _sats_to_msats(settings.min_payout_sat) + ) + if available_balance > min_amount: + amount_received = await raw_send_to_lnurl( + wallet, + proofs, + settings.receive_ln_address, + unit, + amount=available_balance, ) - if available_balance > min_amount: - amount_received = await raw_send_to_lnurl( - wallet, - proofs, - settings.receive_ln_address, - unit, - amount=available_balance, - ) - logger.info( - "Payout sent successfully", - extra={ - "mint_url": mint_url, - "unit": unit, - "balance": available_balance, - "amount_received": amount_received, - }, - ) - except Exception as e: - logger.error( - f"Error sending payout: {type(e).__name__}", + logger.info( + "Payout sent successfully", extra={ - "error": str(e), "mint_url": mint_url, "unit": unit, + "balance": available_balance, + "amount_received": amount_received, }, ) + except Exception as e: + logger.error( + f"Error sending payout: {type(e).__name__}", + extra={ + "error": str(e), + "mint_url": mint_url, + "unit": unit, + }, + ) except Exception as e: logger.error( f"Error in periodic payout cycle: {type(e).__name__}", @@ -993,22 +1043,65 @@ async def periodic_payout() -> None: ) +async def _set_refund_sweep_state( + refund_id: str, + *, + predicates: tuple[typing.Any, ...] = (), + **values: object, +) -> int: + async with db.create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(db.CashuTransaction) + .where(col(db.CashuTransaction.id) == refund_id, *predicates) + .values(**values) + ) + await session.commit() + return int(result.rowcount or 0) + + async def _refund_sweep_once(cutoff: int) -> None: + claim_cutoff = int(time.time()) - settings.refund_sweep_claim_timeout_seconds + claim_available = col(db.CashuTransaction.sweep_started_at).is_(None) | ( + col(db.CashuTransaction.sweep_started_at) < claim_cutoff + ) async with db.create_session() as session: stmt = select(db.CashuTransaction).where( db.CashuTransaction.type == "out", db.CashuTransaction.collected == False, # noqa: E712 db.CashuTransaction.swept == False, # noqa: E712 db.CashuTransaction.created_at < cutoff, + claim_available, ) results = await session.exec(stmt) refunds = results.all() - for refund in refunds: - try: - await recieve_token(refund.token) - refund.swept = True - session.add(refund) + for refund in refunds: + reclaimed_stale_claim = refund.sweep_started_at is not None + claim_started_at = int(time.time()) + claimed = await _set_refund_sweep_state( + refund.id, + predicates=( + col(db.CashuTransaction.swept) == False, # noqa: E712 + col(db.CashuTransaction.collected) == False, # noqa: E712 + claim_available, + ), + sweep_started_at=claim_started_at, + ) + if claimed != 1: + continue + + claim_owned = col(db.CashuTransaction.sweep_started_at) == claim_started_at + redeemed = False + try: + await recieve_token(refund.token) + redeemed = True + finalized = await _set_refund_sweep_state( + refund.id, + predicates=(claim_owned,), + swept=True, + sweep_started_at=None, + ) + if finalized == 1: logger.info( "Swept uncollected refund", extra={ @@ -1017,26 +1110,70 @@ async def _refund_sweep_once(cutoff: int) -> None: "unit": refund.unit, }, ) - except Exception as e: - error_msg = str(e).lower() - if "already spent" in error_msg: - refund.collected = True - session.add(refund) + else: + logger.critical( + "Refund token swept after claim ownership changed; manual reconciliation required", + extra={"id": refund.id}, + ) + except BaseException as e: + if redeemed or isinstance(e, TokenConsumedError): + # The token was spent, or the redemption outcome is known to be + # post-spend. Retain the claim so a stale retry classifies + # "already spent" as swept, never as a client collection. + logger.critical( + "Refund token spent but sweep checkpoint was not completed; manual reconciliation required", + extra={"id": refund.id}, + exc_info=isinstance(e, Exception), + ) + if not isinstance(e, Exception): + raise + continue + + error_msg = str(e).lower() + if isinstance(e, Exception) and "already spent" in error_msg: + if reclaimed_stale_claim: + # A prior worker may have redeemed the token and crashed + # before finalizing. Treat the ambiguous stale claim as a + # completed sweep rather than misreporting client collection. + updated = await _set_refund_sweep_state( + refund.id, + predicates=(claim_owned,), + swept=True, + sweep_started_at=None, + ) + else: + updated = await _set_refund_sweep_state( + refund.id, + predicates=(claim_owned,), + collected=True, + swept=False, + sweep_started_at=None, + ) + if updated == 1: logger.info( - "Refund already spent (client collected), marking swept", + "Refund token was already spent", extra={ "id": refund.id, + "reclaimed_stale_claim": reclaimed_stale_claim, }, ) else: logger.warning( - "Failed to sweep refund", - extra={ - "id": refund.id, - "error": str(e), - }, + "Refund claim ownership changed before spent-token checkpoint", + extra={"id": refund.id}, ) - await session.commit() + else: + # Once redemption starts, an exception cannot prove the token + # was not spent (for example, a melt may land before the + # response is lost). Retain the claim so a stale retry treats + # an "already spent" result as a completed sweep. + logger.critical( + "Refund token redemption outcome is unknown; retaining sweep claim for reconciliation", + extra={"id": refund.id, "error": str(e)}, + exc_info=isinstance(e, Exception), + ) + if not isinstance(e, Exception): + raise async def refund_sweep_once() -> None: @@ -1082,53 +1219,71 @@ async def periodic_routstr_fee_payout() -> None: ) continue - accumulated_sats = fee.accumulated_msats // 1000 - if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT: - wallet = await get_wallet(settings.primary_mint, "sat") - proofs = get_proofs_per_mint_and_unit( - wallet, settings.primary_mint, "sat", not_reserved=True - ) - paid_msats = accumulated_sats * 1000 - payout_checkpointed = await db.reset_routstr_fee( - session, paid_msats - ) - if not payout_checkpointed: - logger.warning("Routstr fee payout was already claimed") - continue + accumulated_sats = _msats_to_sats(fee.accumulated_msats) + if accumulated_sats < ROUTSTR_FEE_DEFAULT_PAYOUT: + continue + paid_msats = _sats_to_msats(accumulated_sats) - try: - amount_received = await raw_send_to_lnurl( - wallet, - proofs, - ROUTSTR_LN_ADDRESS, - "sat", - amount=accumulated_sats, - ) - except Exception: - logger.critical( - "Routstr fee payout outcome is unknown; manual reconciliation required", - extra={"payout_in_progress_msats": paid_msats}, - exc_info=True, - ) - continue + # 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, + ) + 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 ) - if not payout_completed: - logger.critical( - "Routstr fee payout sent but checkpoint was not completed", - extra={"payout_in_progress_msats": paid_msats}, - ) - continue + 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, - }, - ) + 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__}", diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index a26b9083..079954ce 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -3,20 +3,23 @@ Covers two things: - The three constraint fields (balance_limit, balance_limit_reset, validity_date) are persisted on LightningInvoice and survive a DB round-trip. -- create_api_key_from_invoice propagates those fields to the created ApiKey, - so the constraints are actually enforced when the key is used. +- The production-path API-key record helper propagates those fields to the + created ApiKey, so the constraints are actually enforced when the key is used. """ from __future__ import annotations +import asyncio import time -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest +from sqlalchemy import inspect +from sqlalchemy.ext.asyncio import AsyncEngine from sqlmodel.ext.asyncio.session import AsyncSession from routstr.core.db import ApiKey, LightningInvoice -from routstr.lightning import create_api_key_from_invoice +from routstr.lightning import _create_api_key_record def _make_invoice(**kwargs: object) -> LightningInvoice: @@ -48,6 +51,7 @@ def mock_wallet_mint() -> object: # Persistence # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_invoice_persists_balance_limit( integration_session: AsyncSession, @@ -92,6 +96,7 @@ async def test_invoice_persists_validity_date( # Propagation to ApiKey # --------------------------------------------------------------------------- + @pytest.mark.asyncio async def test_created_key_receives_balance_limit( integration_session: AsyncSession, @@ -100,7 +105,7 @@ async def test_created_key_receives_balance_limit( integration_session.add(invoice) await integration_session.flush() - api_key = await create_api_key_from_invoice(invoice, integration_session) + api_key = await _create_api_key_record(invoice, integration_session) await integration_session.commit() stored_key = await integration_session.get(ApiKey, api_key.hashed_key) @@ -116,7 +121,7 @@ async def test_created_key_receives_balance_limit_reset( integration_session.add(invoice) await integration_session.flush() - api_key = await create_api_key_from_invoice(invoice, integration_session) + api_key = await _create_api_key_record(invoice, integration_session) await integration_session.commit() stored_key = await integration_session.get(ApiKey, api_key.hashed_key) @@ -133,7 +138,7 @@ async def test_created_key_receives_validity_date( integration_session.add(invoice) await integration_session.flush() - api_key = await create_api_key_from_invoice(invoice, integration_session) + api_key = await _create_api_key_record(invoice, integration_session) await integration_session.commit() stored_key = await integration_session.get(ApiKey, api_key.hashed_key) @@ -141,6 +146,243 @@ async def test_created_key_receives_validity_date( assert stored_key.validity_date == expiry +@pytest.mark.asyncio +async def test_payment_check_releases_connection_during_mint_quote( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _make_invoice(id="inv_slow_quote", status="pending", paid_at=None) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add(invoice) + await setup.commit() + + async with AsyncSession(integration_engine, expire_on_commit=False) as session: + stored = await session.get(LightningInvoice, invoice.id) + assert stored is not None + + async def quote_status(*args: object, **kwargs: object) -> MagicMock: + assert integration_engine.pool.checkedout() == 0 # type: ignore[attr-defined] + return MagicMock(paid=False) + + wallet = MagicMock() + wallet.get_mint_quote = AsyncMock(side_effect=quote_status) + with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)): + from routstr.lightning import check_invoice_payment + + await check_invoice_payment(stored, session) + + +@pytest.mark.asyncio +async def test_concurrent_payment_checks_mint_and_credit_invoice_once( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _make_invoice(id="inv_concurrent", status="pending", paid_at=None) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add(invoice) + await setup.commit() + + wallet = MagicMock() + wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) + + mint_calls = 0 + + async def single_use_mint(*args: object, **kwargs: object) -> list[object]: + # Real mints enforce single-use quotes: the second concurrent minter + # gets rejected at the mint, mirroring cashu quote semantics. + nonlocal mint_calls + mint_calls += 1 + call_number = mint_calls + await asyncio.sleep(0.05) + if call_number > 1: + raise Exception("quote already issued") + return [] + + wallet.mint = AsyncMock(side_effect=single_use_mint) + + async with ( + AsyncSession(integration_engine, expire_on_commit=False) as first, + AsyncSession(integration_engine, expire_on_commit=False) as second, + ): + first_invoice = await first.get(LightningInvoice, invoice.id) + second_invoice = await second.get(LightningInvoice, invoice.id) + assert first_invoice is not None + assert second_invoice is not None + + with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)): + from routstr.lightning import check_invoice_payment + + await asyncio.gather( + check_invoice_payment(first_invoice, first), + check_invoice_payment(second_invoice, second), + ) + + 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 + assert stored_invoice.status == "paid" + assert stored_invoice.api_key_hash is not None + stored_key = await verify.get(ApiKey, stored_invoice.api_key_hash) + assert stored_key is not None + assert stored_key.balance == invoice.amount_sats * 1000 + + +@pytest.mark.asyncio +async def test_failed_mint_keeps_invoice_pending_for_retry( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _make_invoice(id="inv_mint_failure", status="pending", paid_at=None) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add(invoice) + await setup.commit() + + wallet = MagicMock() + 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: + stored = await session.get(LightningInvoice, invoice.id) + assert stored is not None + with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)): + from routstr.lightning import check_invoice_payment + + await check_invoice_payment(stored, session) + + 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" + + +@pytest.mark.asyncio +async def test_unpaid_topup_does_not_query_target_key( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _make_invoice( + id="inv_unpaid_topup", + status="pending", + paid_at=None, + purpose="topup", + api_key_hash="target-key", + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add(invoice) + await setup.commit() + + wallet = MagicMock() + wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=False)) + create_session = MagicMock(side_effect=RuntimeError("target lookup should not run")) + async with AsyncSession(integration_engine, expire_on_commit=False) as session: + stored = await session.get(LightningInvoice, invoice.id) + assert stored is not None + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning.create_session", create_session), + ): + from routstr.lightning import check_invoice_payment + + await check_invoice_payment(stored, session) + + wallet.get_mint_quote.assert_awaited_once_with(invoice.payment_hash) + create_session.assert_not_called() + + +@pytest.mark.asyncio +async def test_missing_topup_target_is_rejected_before_mint( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _make_invoice( + id="inv_missing_topup_target", + status="pending", + paid_at=None, + purpose="topup", + api_key_hash="pruned-key", + expires_at=int(time.time()) - 1, + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add(invoice) + await setup.commit() + + wallet = MagicMock() + wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) + wallet.mint = AsyncMock() + async with AsyncSession(integration_engine, expire_on_commit=False) as session: + stored = await session.get(LightningInvoice, invoice.id) + assert stored is not None + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning.logger.critical") as critical, + ): + from routstr.lightning import get_invoice_status + + response = await get_invoice_status(invoice.id, session) + + assert response.status == "reconciliation_required" + assert stored.status == "reconciliation_required" + assert stored not in session.dirty + critical.assert_called_once() + + wallet.mint.assert_not_awaited() + async with AsyncSession(integration_engine) as verify: + stored = await verify.get(LightningInvoice, invoice.id) + assert stored is not None + assert stored.status == "reconciliation_required" + + +@pytest.mark.asyncio +async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + invoice = _make_invoice(id="inv_finalize_failure", status="pending", paid_at=None) + sibling = _make_invoice( + id="inv_finalize_failure_sibling", + bolt11="lnbc1000n1sibling", + payment_hash="cafebabe" * 8, + status="pending", + paid_at=None, + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add_all([invoice, sibling]) + await setup.commit() + + wallet = MagicMock() + wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) + wallet.mint = AsyncMock(return_value=[]) + 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) + assert stored is not None + assert stored_sibling is not None + with ( + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch( + "routstr.lightning._create_api_key_record", + AsyncMock(side_effect=RuntimeError("database unavailable")), + ), + ): + from routstr.lightning import check_invoice_payment + + await check_invoice_payment(stored, session) + + stored_state = inspect(stored) + sibling_state = inspect(stored_sibling) + assert stored_state is not None + assert sibling_state is not None + assert stored_state.expired is False + assert sibling_state.expired is False + assert stored.status == "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" + + @pytest.mark.asyncio async def test_created_key_without_constraints_has_none_fields( integration_session: AsyncSession, @@ -149,7 +391,7 @@ async def test_created_key_without_constraints_has_none_fields( integration_session.add(invoice) await integration_session.flush() - api_key = await create_api_key_from_invoice(invoice, integration_session) + api_key = await _create_api_key_record(invoice, integration_session) await integration_session.commit() stored_key = await integration_session.get(ApiKey, api_key.hashed_key) @@ -157,3 +399,84 @@ async def test_created_key_without_constraints_has_none_fields( assert stored_key.balance_limit is None assert stored_key.balance_limit_reset is None assert stored_key.validity_date is None + + +@pytest.mark.asyncio +async def test_db_guard_credits_once_when_both_mints_succeed( + integration_engine: AsyncEngine, + patched_db_engine: None, +) -> None: + """Even if the mint fails to enforce single-use quotes and both racers + mint successfully, the conditional status update must credit exactly once.""" + key = ApiKey(hashed_key="race-key", balance=1_000) + invoice = _make_invoice( + id="inv_db_guard", + status="pending", + paid_at=None, + purpose="topup", + api_key_hash="race-key", + ) + sibling = _make_invoice( + id="inv_db_guard_sibling", + bolt11="lnbc1000n1race-sibling", + payment_hash="01234567" * 8, + status="pending", + paid_at=None, + ) + async with AsyncSession(integration_engine, expire_on_commit=False) as setup: + setup.add_all([key, invoice, sibling]) + await setup.commit() + + wallet = MagicMock() + wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True)) + + async def always_succeeding_mint(*args: object, **kwargs: object) -> list[object]: + await asyncio.sleep(0.05) + return [] + + wallet.mint = AsyncMock(side_effect=always_succeeding_mint) + + async with ( + AsyncSession(integration_engine, expire_on_commit=False) as first, + AsyncSession(integration_engine, expire_on_commit=False) as second, + ): + first_invoice = await first.get(LightningInvoice, invoice.id) + first_sibling = await first.get(LightningInvoice, sibling.id) + second_invoice = await second.get(LightningInvoice, invoice.id) + assert first_invoice is not None + assert first_sibling is not None + assert second_invoice is not None + + with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)): + from routstr.lightning import check_invoice_payment + + await asyncio.gather( + check_invoice_payment(first_invoice, first), + check_invoice_payment(second_invoice, second), + ) + + first_state = inspect(first_invoice) + sibling_state = inspect(first_sibling) + second_state = inspect(second_invoice) + assert first_state is not None + assert sibling_state is not None + assert second_state is not None + assert first_state.expired is False + assert sibling_state.expired is False + assert second_state.expired is False + assert first_invoice.id == invoice.id + assert first_sibling.id == sibling.id + assert second_invoice.id == invoice.id + assert first_invoice.status == "paid" + assert second_invoice.status == "paid" + assert first_invoice not in first.dirty + assert second_invoice not in second.dirty + + assert wallet.mint.await_count == 2 + 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 + assert stored_invoice.status == "paid" + stored_key = await verify.get(ApiKey, "race-key") + assert stored_key is not None + assert stored_key.balance == 1_000 + invoice.amount_sats * 1000 diff --git a/tests/unit/test_admin_transactions.py b/tests/unit/test_admin_transactions.py new file mode 100644 index 00000000..22508fd6 --- /dev/null +++ b/tests/unit/test_admin_transactions.py @@ -0,0 +1,35 @@ +from contextlib import asynccontextmanager +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from routstr.core.admin import get_transactions_api +from routstr.core.db import CashuTransaction + + +@pytest.mark.asyncio +async def test_transactions_api_excludes_internal_sweep_claim_timestamp() -> None: + transaction = CashuTransaction( + token="cashu-token", + amount=10, + unit="sat", + type="out", + sweep_started_at=123, + ) + count_result = MagicMock() + count_result.one.return_value = 1 + transactions_result = MagicMock() + transactions_result.all.return_value = [transaction] + session = MagicMock() + session.exec = AsyncMock(side_effect=[count_result, transactions_result]) + + @asynccontextmanager + async def create_session(): # type: ignore[no-untyped-def] + yield session + + with patch("routstr.core.admin.create_session", create_session): + response = await get_transactions_api() + + assert response["total"] == 1 + assert response["transactions"][0]["token"] == "cashu-token" + assert "sweep_started_at" not in response["transactions"][0] diff --git a/tests/unit/test_balances_by_mint_and_unit.py b/tests/unit/test_balances_by_mint_and_unit.py new file mode 100644 index 00000000..248c43a8 --- /dev/null +++ b/tests/unit/test_balances_by_mint_and_unit.py @@ -0,0 +1,115 @@ +"""Real-DB coverage for db.balances_by_mint_and_unit. + +Verifies the grouped liability query used by fetch_all_balances: it sums +balances per (mint_url, unit), filters to the requested mints/units, excludes +NULL mint/currency rows, and returns nothing for empty inputs. +""" + +from typing import AsyncGenerator + +import pytest +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlalchemy.pool import StaticPool +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ( + ApiKey, + balance_for_mint_and_unit, + balances_by_mint_and_unit, +) + + +def _make_engine() -> AsyncEngine: + return create_async_engine( + "sqlite+aiosqlite://", + poolclass=StaticPool, + connect_args={"check_same_thread": False}, + ) + + +@pytest.fixture +async def session() -> "AsyncGenerator[AsyncSession, None]": + engine = _make_engine() + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + db_session = AsyncSession(engine, expire_on_commit=False) + try: + yield db_session + finally: + await db_session.close() + await engine.dispose() + + +async def _add_key( + session: AsyncSession, + hashed_key: str, + balance: int, + mint_url: str | None, + currency: str | None, +) -> None: + session.add( + ApiKey( + hashed_key=hashed_key, + balance=balance, + refund_mint_url=mint_url, + refund_currency=currency, + ) + ) + await session.commit() + + +@pytest.mark.asyncio +async def test_sums_and_groups_by_mint_and_unit(session: AsyncSession) -> None: + await _add_key(session, "a", 1000, "http://m1", "sat") + await _add_key(session, "b", 500, "http://m1", "sat") + await _add_key(session, "c", 7000, "http://m1", "msat") + await _add_key(session, "d", 200, "http://m2", "sat") + + result = await balances_by_mint_and_unit( + session, ["http://m1", "http://m2"], ["sat", "msat"] + ) + + assert result[("http://m1", "sat")] == 1500 + assert result[("http://m1", "msat")] == 7000 + assert result[("http://m2", "sat")] == 200 + + +@pytest.mark.asyncio +async def test_filters_out_unrequested_mints_and_units(session: AsyncSession) -> None: + await _add_key(session, "a", 1000, "http://wanted", "sat") + await _add_key(session, "b", 999, "http://other", "sat") + await _add_key(session, "c", 888, "http://wanted", "usd") + + result = await balances_by_mint_and_unit(session, ["http://wanted"], ["sat"]) + + assert result == {("http://wanted", "sat"): 1000} + + +@pytest.mark.asyncio +async def test_excludes_rows_with_null_mint_or_currency(session: AsyncSession) -> None: + await _add_key(session, "a", 1000, "http://m1", "sat") + await _add_key(session, "b", 4242, None, None) + + result = await balances_by_mint_and_unit(session, ["http://m1"], ["sat"]) + + assert result == {("http://m1", "sat"): 1000} + + +@pytest.mark.asyncio +async def test_scalar_balance_for_one_mint_and_unit(session: AsyncSession) -> None: + await _add_key(session, "a", 1000, "http://m1", "sat") + await _add_key(session, "b", 500, "http://m1", "sat") + await _add_key(session, "c", 9000, "http://m1", "msat") + await _add_key(session, "d", 700, "http://m2", "sat") + + assert await balance_for_mint_and_unit(session, "http://m1", "sat") == 1500 + assert await balance_for_mint_and_unit(session, "http://missing", "sat") == 0 + + +@pytest.mark.asyncio +async def test_empty_inputs_return_empty_mapping(session: AsyncSession) -> None: + await _add_key(session, "a", 1000, "http://m1", "sat") + + assert await balances_by_mint_and_unit(session, [], ["sat"]) == {} + assert await balances_by_mint_and_unit(session, ["http://m1"], []) == {} diff --git a/tests/unit/test_db_pool_config.py b/tests/unit/test_db_pool_config.py new file mode 100644 index 00000000..8f5d0367 --- /dev/null +++ b/tests/unit/test_db_pool_config.py @@ -0,0 +1,85 @@ +from unittest.mock import MagicMock, patch + +import pytest +from sqlalchemy.pool import StaticPool + +from routstr.core import db +from routstr.core.db import create_db_engine +from routstr.core.settings import settings + + +@pytest.mark.asyncio +async def test_engine_uses_validated_bounded_pool_settings( + monkeypatch: pytest.MonkeyPatch, tmp_path: object +) -> None: + monkeypatch.setattr(settings, "database_pool_size", 12) + monkeypatch.setattr(settings, "database_max_overflow", 3) + monkeypatch.setattr(settings, "database_pool_timeout", 2.5) + monkeypatch.setattr(settings, "database_pool_recycle", 900) + monkeypatch.setattr(settings, "database_pool_pre_ping", False) + + engine = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/pool.db") + try: + assert engine.pool.size() == 12 # type: ignore[attr-defined] + assert engine.pool._max_overflow == 3 # type: ignore[attr-defined] + assert engine.pool._timeout == 2.5 # type: ignore[attr-defined] + assert engine.pool._recycle == 900 + assert engine.pool._pre_ping is False + finally: + await engine.dispose() + + +@pytest.mark.asyncio +async def test_memory_sqlite_keeps_static_pool( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "database_pool_pre_ping", True) + engine = create_db_engine("sqlite+aiosqlite://") + try: + assert isinstance(engine.pool, StaticPool) + assert engine.pool._pre_ping is True + finally: + await engine.dispose() + + +def test_non_sqlite_backend_enables_pre_ping_automatically( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "database_pool_pre_ping", False) + fake_engine = MagicMock() + + with ( + patch.object(db, "create_async_engine", return_value=fake_engine) as factory, + patch.object(db.event, "listen") as listen, + ): + created = create_db_engine("postgresql+asyncpg://user:pass@db/node") + + assert created is fake_engine + assert factory.call_args.kwargs["pool_pre_ping"] is True + assert listen.call_count == 2 + + +@pytest.mark.asyncio +async def test_every_created_engine_warns_for_long_checkouts( + monkeypatch: pytest.MonkeyPatch, tmp_path: object +) -> None: + monkeypatch.setattr(settings, "database_pool_hold_warn_seconds", 0.0) + monkeypatch.setattr(settings, "database_pool_pre_ping", False) + first = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/first.db") + second = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/second.db") + + try: + with patch.object(db.logger, "warning") as warning: + async with first.connect() as connection: + await connection.exec_driver_sql("SELECT 1") + async with second.connect() as connection: + await connection.exec_driver_sql("SELECT 1") + + assert warning.call_count == 2 + assert all( + call.kwargs["extra"]["threshold_seconds"] == 0.0 + for call in warning.call_args_list + ) + finally: + await first.dispose() + await second.dispose() diff --git a/tests/unit/test_fee_payout_crash_safety.py b/tests/unit/test_fee_payout_crash_safety.py index 0939baa4..ae7d0788 100644 --- a/tests/unit/test_fee_payout_crash_safety.py +++ b/tests/unit/test_fee_payout_crash_safety.py @@ -1,5 +1,5 @@ import asyncio -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch @@ -13,9 +13,19 @@ from routstr import wallet from routstr.core import db -@asynccontextmanager -async def _session_context(session: Mock) -> AsyncIterator[Mock]: - yield session +class _SessionContext: + def __init__(self, session: Mock) -> None: + self.session = session + + async def __aenter__(self) -> Mock: + return self.session + + async def __aexit__(self, *args: object) -> None: + return None + + +def _session_context(session: Mock) -> _SessionContext: + return _SessionContext(session) @pytest.mark.asyncio @@ -47,7 +57,7 @@ async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None: @pytest.mark.asyncio -async def test_fee_payout_checkpoints_before_sending() -> None: +async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> None: session = Mock() fee = SimpleNamespace( accumulated_msats=5_000, @@ -57,6 +67,10 @@ async def test_fee_payout_checkpoints_before_sending() -> None: payout_wallet = Mock() events: list[str] = [] + async def prepare(*_args: object) -> Mock: + events.append("prepare") + return payout_wallet + async def checkpoint(*_args: object) -> bool: events.append("checkpoint") return True @@ -77,18 +91,92 @@ async def test_fee_payout_checkpoints_before_sending() -> None: "routstr.wallet.asyncio.sleep", AsyncMock(side_effect=[None, asyncio.CancelledError()]), ), - patch("routstr.wallet.db.create_session", return_value=_session_context(session)), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint), patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete), - patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)), + patch("routstr.wallet.get_wallet", AsyncMock(side_effect=prepare)), patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), patch("routstr.wallet.raw_send_to_lnurl", side_effect=send), ): with pytest.raises(asyncio.CancelledError): await wallet.periodic_routstr_fee_payout() - assert events == ["checkpoint", "send", "complete"] + assert events == ["prepare", "checkpoint", "send", "complete"] + + +@pytest.mark.asyncio +async def test_fee_payout_preparation_failure_does_not_checkpoint() -> None: + session = Mock() + fee = SimpleNamespace( + accumulated_msats=5_000, + payout_in_progress_msats=0, + payout_started_at=None, + ) + checkpoint = AsyncMock() + + with ( + patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1), + patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1), + patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"), + patch( + "routstr.wallet.asyncio.sleep", + AsyncMock(side_effect=[None, asyncio.CancelledError()]), + ), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), + patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), + patch("routstr.wallet.db.reset_routstr_fee", checkpoint), + patch( + "routstr.wallet.get_wallet", + AsyncMock(side_effect=RuntimeError("wallet unavailable")), + ), + ): + with pytest.raises(asyncio.CancelledError): + await wallet.periodic_routstr_fee_payout() + + checkpoint.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None: + session = Mock() + fee = SimpleNamespace( + accumulated_msats=5_000, + payout_in_progress_msats=0, + payout_started_at=None, + ) + send = AsyncMock() + + with ( + patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1), + patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1), + patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"), + patch( + "routstr.wallet.asyncio.sleep", + AsyncMock(side_effect=[None, asyncio.CancelledError()]), + ), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), + patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), + patch( + "routstr.wallet.db.reset_routstr_fee", + AsyncMock(return_value=False), + ), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())), + patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), + patch("routstr.wallet.raw_send_to_lnurl", send), + patch("routstr.wallet.logger.warning") as warning, + ): + with pytest.raises(asyncio.CancelledError): + await wallet.periodic_routstr_fee_payout() + + send.assert_not_awaited() + warning.assert_called_once_with("Routstr fee payout was already claimed") @pytest.mark.asyncio @@ -107,7 +195,9 @@ async def test_fee_payout_does_not_retry_an_unresolved_checkpoint() -> None: "routstr.wallet.asyncio.sleep", AsyncMock(side_effect=[None, asyncio.CancelledError()]), ), - patch("routstr.wallet.db.create_session", return_value=_session_context(session)), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), patch("routstr.wallet.db.reset_routstr_fee", AsyncMock()) as checkpoint, patch("routstr.wallet.get_wallet", AsyncMock()) as get_wallet, @@ -141,7 +231,9 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non "routstr.wallet.asyncio.sleep", AsyncMock(side_effect=[None, asyncio.CancelledError()]), ), - patch("routstr.wallet.db.create_session", return_value=_session_context(session)), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)), patch("routstr.wallet.db.complete_routstr_fee_payout", complete), @@ -158,3 +250,139 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non complete.assert_not_awaited() critical.assert_called_once() + + +@pytest.mark.asyncio +async def test_fee_payout_cancellation_during_send_alerts_and_propagates() -> None: + session = Mock() + fee = SimpleNamespace( + accumulated_msats=5_000, + payout_in_progress_msats=0, + payout_started_at=None, + ) + complete = AsyncMock() + + with ( + patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1), + patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1), + patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"), + patch("routstr.wallet.asyncio.sleep", AsyncMock(return_value=None)), + patch( + "routstr.wallet.db.create_session", return_value=_session_context(session) + ), + patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), + patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)), + patch("routstr.wallet.db.complete_routstr_fee_payout", complete), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())), + patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), + patch( + "routstr.wallet.raw_send_to_lnurl", + AsyncMock(side_effect=asyncio.CancelledError()), + ), + patch("routstr.wallet.logger.critical") as critical, + ): + with pytest.raises(asyncio.CancelledError): + await wallet.periodic_routstr_fee_payout() + + complete.assert_not_awaited() + critical.assert_called_once() + assert critical.call_args.args[0] == ( + "Routstr fee payout outcome is unknown; manual reconciliation required" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure_site", ["session", "completion"]) +async def test_fee_payout_completion_failures_use_sent_checkpoint_alert( + failure_site: str, +) -> None: + session = Mock() + fee = SimpleNamespace( + accumulated_msats=5_000, + payout_in_progress_msats=0, + payout_started_at=None, + ) + completion = AsyncMock() + if failure_site == "session": + create_session = Mock( + side_effect=[ + _session_context(session), + _session_context(session), + RuntimeError("pool unavailable"), + ] + ) + else: + create_session = Mock(return_value=_session_context(session)) + completion.side_effect = RuntimeError("checkpoint unavailable") + + with ( + patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1), + patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1), + patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"), + patch( + "routstr.wallet.asyncio.sleep", + AsyncMock(side_effect=[None, asyncio.CancelledError()]), + ), + patch("routstr.wallet.db.create_session", create_session), + patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), + patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)), + patch("routstr.wallet.db.complete_routstr_fee_payout", completion), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())), + patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), + patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(return_value=5)), + patch("routstr.wallet.logger.critical") as critical, + ): + with pytest.raises(asyncio.CancelledError): + await wallet.periodic_routstr_fee_payout() + + critical.assert_called_once() + assert critical.call_args.args[0] == ( + "Routstr fee payout sent but checkpoint was not completed" + ) + + +@pytest.mark.asyncio +async def test_fee_payout_releases_db_connection_during_send(tmp_path: object) -> None: + """With pool_size=1, the payout must not hold a connection while the + external LNURL send is in flight, or the completion step would starve.""" + engine = create_async_engine( + f"sqlite+aiosqlite:///{tmp_path}/payout.db", pool_size=1, max_overflow=0 + ) + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + async with AsyncSession(engine) as session: + session.add(db.RoutstrFee(id=1, accumulated_msats=5_000_000)) + await session.commit() + + @asynccontextmanager + async def create_session() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + async def send(*_args: object, **_kwargs: object) -> int: + assert engine.pool.checkedout() == 0 # type: ignore[attr-defined] + return 5 + + try: + with ( + patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1), + patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1), + patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"), + patch( + "routstr.wallet.asyncio.sleep", + AsyncMock(side_effect=[None, asyncio.CancelledError()]), + ), + patch("routstr.wallet.db.create_session", create_session), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())), + patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), + patch("routstr.wallet.raw_send_to_lnurl", side_effect=send), + ): + with pytest.raises(asyncio.CancelledError): + await wallet.periodic_routstr_fee_payout() + + async with AsyncSession(engine) as session: + fee = await db.get_routstr_fee(session) + assert fee.payout_in_progress_msats == 0 + assert fee.total_paid_msats == 5_000_000 + finally: + await engine.dispose() diff --git a/tests/unit/test_fee_payout_migration.py b/tests/unit/test_fee_payout_migration.py index d820a056..8eae9829 100644 --- a/tests/unit/test_fee_payout_migration.py +++ b/tests/unit/test_fee_payout_migration.py @@ -37,7 +37,7 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None: "payout_in_progress_msats, payout_started_at FROM routstr_fees" ).fetchone() - assert version == ("9c4d8e2f1a6b",) + assert version == ("aa50fde387a2",) assert { "id", "accumulated_msats", diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index 433e84d4..3427cf78 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -1,7 +1,13 @@ +import asyncio +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager +from pathlib import Path from unittest.mock import AsyncMock, MagicMock, patch import pytest +from sqlalchemy.ext.asyncio import create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession from routstr.wallet import fetch_all_balances @@ -26,8 +32,10 @@ def _patches( # type: ignore[no-untyped-def] AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch( - "routstr.wallet.db.balances_for_mint_and_unit", - AsyncMock(return_value=user_balance_msats), + "routstr.wallet.db.balances_by_mint_and_unit", + AsyncMock( + return_value={("http://primary:3338", "sat"): user_balance_msats} + ), ), patch("routstr.wallet.db.create_session", _fake_session), ] @@ -38,8 +46,9 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None: """With empty cashu_mints, balances are still fetched for primary_mint.""" from routstr.core.settings import settings - with patch.object(settings, "cashu_mints", []), patch.object( - settings, "primary_mint", "http://primary:3338" + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), ): for p in _patches(proof_amount=1000): p.start() @@ -54,13 +63,159 @@ 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_closes_db_session_before_concurrent_mint_io() -> None: + """Slow mint checks must never run while the balance DB session is open.""" + from routstr.core.settings import settings + + session_open = False + mint_calls = 0 + + @asynccontextmanager + async def tracked_session(): # type: ignore[no-untyped-def] + nonlocal session_open + session_open = True + try: + yield MagicMock() + finally: + session_open = False + + async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def] + nonlocal mint_calls + assert session_open is False + mint_calls += 1 + await asyncio.sleep(0) + return proofs + + with ( + patch.object(settings, "cashu_mints", ["http://one:3338", "http://two:3338"]), + patch.object(settings, "primary_mint", "http://one:3338"), + patch("routstr.wallet.db.create_session", tracked_session), + patch( + "routstr.wallet.db.balances_by_mint_and_unit", + AsyncMock(return_value={}), + create=True, + ), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[MagicMock(amount=1)]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=slow_filter), + ), + ): + details, *_ = await fetch_all_balances(units=["sat", "msat"]) + + assert mint_calls == 4 + assert all("error" not in detail for detail in details) + + +@pytest.mark.asyncio +async def test_fetch_all_balances_bounds_parallel_mint_checks() -> None: + """A slow mint fleet cannot create an unbounded external-I/O fan-out.""" + from routstr.core.settings import settings + + active = 0 + peak = 0 + + async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def] + nonlocal active, peak + active += 1 + peak = max(peak, active) + await asyncio.sleep(0.01) + active -= 1 + return proofs + + with ( + patch.object( + settings, + "cashu_mints", + [f"http://mint-{index}:3338" for index in range(8)], + ), + patch.object(settings, "primary_mint", ""), + patch.object(settings, "mint_operation_concurrency", 2), + patch("routstr.wallet.db.create_session", _fake_session), + patch( + "routstr.wallet.db.balances_by_mint_and_unit", + AsyncMock(return_value={}), + ), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=slow_filter), + ), + ): + details, *_ = await fetch_all_balances(units=["sat"]) + + assert len(details) == 8 + assert peak == 2 + + +@pytest.mark.asyncio +async def test_slow_mints_do_not_exhaust_a_single_connection_pool( + tmp_path: Path, +) -> None: + """Concurrent slow balance refreshes release the sole DB connection promptly.""" + from routstr.core.settings import settings + + engine = create_async_engine( + f"sqlite+aiosqlite:///{tmp_path / 'pool-pressure.db'}", + pool_size=1, + max_overflow=0, + pool_timeout=0.2, + ) + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + + @asynccontextmanager + async def single_pool_session() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def] + await asyncio.sleep(0.3) + return proofs + + try: + with ( + patch.object(settings, "cashu_mints", ["http://slow:3338"]), + patch.object(settings, "primary_mint", "http://slow:3338"), + patch.object(settings, "mint_operation_concurrency", 1), + patch("routstr.wallet.db.create_session", single_pool_session), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=slow_filter), + ), + ): + results = await asyncio.gather( + *(fetch_all_balances(units=["sat"]) for _ in range(6)) + ) + + assert all("error" not in result[0][0] for result in results) + assert engine.pool.checkedout() == 0 # type: ignore[attr-defined] + finally: + await engine.dispose() + + @pytest.mark.asyncio async def test_fetch_all_balances_reports_liability_when_wallet_is_empty() -> None: """An empty wallet must not hide outstanding user liabilities.""" from routstr.core.settings import settings - with patch.object(settings, "cashu_mints", []), patch.object( - settings, "primary_mint", "http://primary:3338" + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), ): for p in _patches(proof_amount=0, user_balance_msats=5000): p.start() @@ -84,9 +239,10 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None: """primary_mint already in cashu_mints is not inspected twice.""" from routstr.core.settings import settings - with patch.object( - settings, "cashu_mints", ["http://primary:3338"] - ), patch.object(settings, "primary_mint", "http://primary:3338"): + with ( + patch.object(settings, "cashu_mints", ["http://primary:3338"]), + patch.object(settings, "primary_mint", "http://primary:3338"), + ): for p in _patches(proof_amount=1000): p.start() try: @@ -98,3 +254,58 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None: assert [d["mint_url"] for d in details] == ["http://primary:3338"] assert total_wallet == 1000 + + +@pytest.mark.asyncio +async def test_fetch_all_balances_degrades_when_liability_read_fails() -> None: + from routstr.core.settings import settings + + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), + patch("routstr.wallet.db.create_session", _fake_session), + patch( + "routstr.wallet.db.balances_by_mint_and_unit", + AsyncMock(side_effect=RuntimeError("db pool exhausted")), + ), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[MagicMock(amount=1000)]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=lambda proofs, wallet: proofs), + ), + ): + details, total_wallet, total_user, owner = await fetch_all_balances( + units=["sat"] + ) + + assert details[0]["error"] == "db pool exhausted" + assert details[0]["wallet_balance"] == 1000 + assert details[0]["user_balance"] == 0 + assert details[0]["owner_balance"] == 0 + assert (total_wallet, total_user, owner) == (1000, 0, 0) + + +@pytest.mark.asyncio +async def test_liability_error_keeps_more_specific_mint_error() -> None: + from routstr.core.settings import settings + + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), + patch("routstr.wallet.db.create_session", _fake_session), + patch( + "routstr.wallet.db.balances_by_mint_and_unit", + AsyncMock(side_effect=RuntimeError("db pool exhausted")), + ), + patch( + "routstr.wallet.get_wallet", + AsyncMock(side_effect=RuntimeError("mint down")), + ), + ): + details, *_ = await fetch_all_balances(units=["sat"]) + + assert details[0]["error"] == "mint down" diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index 54ebadb4..953f06cb 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -59,23 +59,29 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non get_wallet = AsyncMock(return_value=MagicMock()) raw_send = AsyncMock(return_value=1000) - with patch.object(settings, "cashu_mints", []), patch.object( - settings, "primary_mint", "http://primary:3338" - ), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object( - settings, "payout_interval_seconds", _INTERVAL - ), patch.object(settings, "min_payout_sat", 10), patch( - "routstr.wallet.asyncio.sleep", _one_cycle_sleep() - ), patch("routstr.wallet.db.create_session", _fake_session), patch( - "routstr.wallet.get_wallet", get_wallet - ), patch( - "routstr.wallet.get_proofs_per_mint_and_unit", - MagicMock(return_value=[MagicMock(amount=100_000)]), - ), patch( - "routstr.wallet.slow_filter_spend_proofs", - AsyncMock(side_effect=lambda proofs, wallet: proofs), - ), patch( - "routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0) - ), patch("routstr.wallet.raw_send_to_lnurl", raw_send): + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), + patch.object(settings, "receive_ln_address", "owner@ln.tld"), + patch.object(settings, "payout_interval_seconds", _INTERVAL), + patch.object(settings, "min_payout_sat", 10), + patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), + patch("routstr.wallet.db.create_session", _fake_session), + patch("routstr.wallet.get_wallet", get_wallet), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[MagicMock(amount=100_000)]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=lambda proofs, wallet: proofs), + ), + patch( + "routstr.wallet.db.balance_for_mint_and_unit", + AsyncMock(return_value=0), + ), + patch("routstr.wallet.raw_send_to_lnurl", raw_send), + ): with pytest.raises(_LoopBreak): await periodic_payout() @@ -84,6 +90,58 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non assert raw_send.await_count >= 1 +@pytest.mark.asyncio +async def test_periodic_payout_releases_session_before_slow_mint_send() -> None: + """The DB connection is returned before the external LNURL call starts.""" + from routstr.core.settings import settings + + session_open = False + sends_completed = 0 + + @asynccontextmanager + async def tracked_session(): # type: ignore[no-untyped-def] + nonlocal session_open + session_open = True + try: + yield MagicMock() + finally: + session_open = False + + async def raw_send(*args: object, **kwargs: object) -> int: + nonlocal sends_completed + assert session_open is False + sends_completed += 1 + return 1000 + + with ( + patch.object(settings, "cashu_mints", []), + patch.object(settings, "primary_mint", "http://primary:3338"), + patch.object(settings, "receive_ln_address", "owner@ln.tld"), + patch.object(settings, "payout_interval_seconds", _INTERVAL), + patch.object(settings, "min_payout_sat", 10), + patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), + patch("routstr.wallet.db.create_session", tracked_session), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[MagicMock(amount=100_000)]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=lambda proofs, wallet: proofs), + ), + patch( + "routstr.wallet.db.balance_for_mint_and_unit", + AsyncMock(return_value=0), + ), + patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)), + ): + with pytest.raises(_LoopBreak): + await periodic_payout() + + assert sends_completed == 2 + + @pytest.mark.asyncio async def test_periodic_payout_isolates_failing_mint() -> None: """A failing mint does not prevent payout for the other mints.""" @@ -97,23 +155,29 @@ async def test_periodic_payout_isolates_failing_mint() -> None: get_wallet = AsyncMock(side_effect=_get_wallet) raw_send = AsyncMock(return_value=1000) - with patch.object( - settings, "cashu_mints", ["http://bad:3338", "http://good:3338"] - ), patch.object(settings, "primary_mint", "http://good:3338"), patch.object( - settings, "receive_ln_address", "owner@ln.tld" - ), patch.object(settings, "payout_interval_seconds", _INTERVAL), patch.object( - settings, "min_payout_sat", 10 - ), patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), patch( - "routstr.wallet.db.create_session", _fake_session - ), patch("routstr.wallet.get_wallet", get_wallet), patch( - "routstr.wallet.get_proofs_per_mint_and_unit", - MagicMock(return_value=[MagicMock(amount=100_000)]), - ), patch( - "routstr.wallet.slow_filter_spend_proofs", - AsyncMock(side_effect=lambda proofs, wallet: proofs), - ), patch( - "routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0) - ), patch("routstr.wallet.raw_send_to_lnurl", raw_send): + with ( + patch.object(settings, "cashu_mints", ["http://bad:3338", "http://good:3338"]), + patch.object(settings, "primary_mint", "http://good:3338"), + patch.object(settings, "receive_ln_address", "owner@ln.tld"), + patch.object(settings, "payout_interval_seconds", _INTERVAL), + patch.object(settings, "min_payout_sat", 10), + patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), + patch("routstr.wallet.db.create_session", _fake_session), + patch("routstr.wallet.get_wallet", get_wallet), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[MagicMock(amount=100_000)]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=lambda proofs, wallet: proofs), + ), + patch( + "routstr.wallet.db.balance_for_mint_and_unit", + AsyncMock(return_value=0), + ), + patch("routstr.wallet.raw_send_to_lnurl", raw_send), + ): with pytest.raises(_LoopBreak): await periodic_payout() @@ -128,27 +192,39 @@ async def test_periodic_payout_isolates_failing_mint() -> None: @pytest.mark.asyncio async def test_periodic_payout_handles_session_creation_failure() -> None: - """A db.create_session failure is logged and the payout loop continues.""" + """A db.create_session failure is logged per mint/unit and the loop continues.""" from routstr.core.settings import settings create_session = MagicMock(side_effect=RuntimeError("db unavailable")) logger = MagicMock() - with patch.object(settings, "cashu_mints", ["http://mint:3338"]), patch.object( - settings, "primary_mint", "http://mint:3338" - ), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object( - settings, "payout_interval_seconds", _INTERVAL - ), patch( - "routstr.wallet.asyncio.sleep", _one_cycle_sleep() - ), patch( - "routstr.wallet.db.create_session", create_session - ), patch("routstr.wallet.logger", logger): + with ( + patch.object(settings, "cashu_mints", ["http://mint:3338"]), + patch.object(settings, "primary_mint", "http://mint:3338"), + patch.object(settings, "receive_ln_address", "owner@ln.tld"), + patch.object(settings, "payout_interval_seconds", _INTERVAL), + patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), + patch("routstr.wallet.db.create_session", create_session), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + MagicMock(return_value=[MagicMock(amount=100_000)]), + ), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=lambda proofs, wallet: proofs), + ), + patch("routstr.wallet.logger", logger), + ): with pytest.raises(_LoopBreak): await periodic_payout() - create_session.assert_called_once() - logger.error.assert_called_once() + # The liability session is opened per mint/unit (sat + msat), and each + # DB failure retains the cycle-specific alert wording while remaining + # isolated to its own iteration. + assert create_session.call_count == 2 + assert logger.error.call_count == 2 message = logger.error.call_args.args[0] extra = logger.error.call_args.kwargs["extra"] assert message == "Error in periodic payout cycle: RuntimeError" - assert extra == {"error": "db unavailable"} + assert extra["error"] == "db unavailable" diff --git a/tests/unit/test_proxy_session_lifecycle.py b/tests/unit/test_proxy_session_lifecycle.py new file mode 100644 index 00000000..5d0416d5 --- /dev/null +++ b/tests/unit/test_proxy_session_lifecycle.py @@ -0,0 +1,43 @@ +from collections.abc import AsyncIterator +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi.responses import StreamingResponse + +from routstr import proxy as proxy_module + + +@pytest.mark.asyncio +async def test_proxy_closes_request_session_before_returning_response() -> None: + """Route completion must release DB resources before response delivery.""" + request = MagicMock() + request.method = "GET" + request.headers = {"accept": "application/json"} + request.url.path = "/not-an-api-route" + request.state.request_id = "test-request" + session = AsyncMock() + + response = await proxy_module.proxy(request, "not-an-api-route", session=session) + + assert response.status_code == 404 + session.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_proxy_session_is_closed_before_first_stream_chunk() -> None: + request = MagicMock() + session = AsyncMock() + + async def stream() -> AsyncIterator[bytes]: + session.close.assert_awaited_once() + yield b"chunk" + + upstream_response = StreamingResponse(stream()) + with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)): + response = await proxy_module.proxy( + request, "v1/chat/completions", session=session + ) + + assert isinstance(response, StreamingResponse) + chunks = [chunk async for chunk in response.body_iterator] + assert chunks == [b"chunk"] diff --git a/tests/unit/test_refund_sweep.py b/tests/unit/test_refund_sweep.py index e82c2111..f2937936 100644 --- a/tests/unit/test_refund_sweep.py +++ b/tests/unit/test_refund_sweep.py @@ -1,4 +1,6 @@ +import asyncio from collections.abc import AsyncIterator +from contextlib import asynccontextmanager from pathlib import Path from unittest.mock import AsyncMock, patch @@ -7,6 +9,7 @@ from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine from sqlmodel import SQLModel, select from sqlmodel.ext.asyncio.session import AsyncSession +from routstr import wallet from routstr.core.db import CashuTransaction from routstr.wallet import refund_sweep_once @@ -39,6 +42,43 @@ async def _load( return {row.token: row for row in result.all()} +@pytest.mark.asyncio +async def test_refund_sweep_releases_db_session_during_token_redemption( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _insert( + session_factory, + CashuTransaction( + token="eligible", amount=1, unit="sat", type="out", created_at=800 + ), + ) + session_open = False + + @asynccontextmanager + async def tracked_session() -> AsyncIterator[AsyncSession]: + nonlocal session_open + async with session_factory() as session: + session_open = True + try: + yield session + finally: + session_open = False + + async def receive_token(token: str) -> None: + assert token == "eligible" + assert session_open is False + + with ( + patch("routstr.wallet.db.create_session", tracked_session), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.time.time", return_value=1000), + patch("routstr.wallet.recieve_token", AsyncMock(side_effect=receive_token)), + ): + await refund_sweep_once() + + assert (await _load(session_factory))["eligible"].swept is True + + @pytest.mark.asyncio async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens( session_factory: async_sessionmaker[AsyncSession], @@ -90,14 +130,17 @@ async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens( @pytest.mark.asyncio @pytest.mark.parametrize( - ("error", "collected"), + ("error", "collected", "claim_started_at"), [ - (RuntimeError("token already spent"), True), - (RuntimeError("mint unavailable"), False), + (RuntimeError("token already spent"), True, None), + (RuntimeError("mint unavailable"), False, 1000), ], ) -async def test_refund_sweep_records_terminal_but_not_transient_failures( - session_factory: async_sessionmaker[AsyncSession], error: Exception, collected: bool +async def test_refund_sweep_records_spent_and_unknown_outcomes_safely( + session_factory: async_sessionmaker[AsyncSession], + error: Exception, + collected: bool, + claim_started_at: int | None, ) -> None: await _insert( session_factory, @@ -116,3 +159,236 @@ async def test_refund_sweep_records_terminal_but_not_transient_failures( refund = (await _load(session_factory))["refund"] assert refund.collected is collected assert refund.swept is False + assert refund.sweep_started_at == claim_started_at + + +@pytest.mark.asyncio +async def test_post_spend_failure_retains_claim_and_stale_retry_records_sweep( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _insert( + session_factory, + CashuTransaction( + token="post-spend-failure", + amount=1, + unit="sat", + type="out", + created_at=800, + ), + ) + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200), + patch("routstr.wallet.time.time", return_value=1000), + patch( + "routstr.wallet.recieve_token", + AsyncMock( + side_effect=wallet.TokenConsumedError( + "Mint on primary failed after successful melt" + ) + ), + ), + ): + await refund_sweep_once() + + retained = (await _load(session_factory))["post-spend-failure"] + assert retained.swept is False + assert retained.collected is False + assert retained.sweep_started_at == 1000 + + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200), + patch("routstr.wallet.time.time", return_value=1300), + patch( + "routstr.wallet.recieve_token", + AsyncMock(side_effect=RuntimeError("token already spent")), + ), + ): + await refund_sweep_once() + + recovered = (await _load(session_factory))["post-spend-failure"] + assert recovered.swept is True + assert recovered.collected is False + assert recovered.sweep_started_at is None + + +@pytest.mark.asyncio +async def test_refund_sweep_retains_claim_on_cancellation_during_redemption( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _insert( + session_factory, + CashuTransaction( + token="cancelled", amount=1, unit="sat", type="out", created_at=800 + ), + ) + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.time.time", return_value=1000), + patch( + "routstr.wallet.recieve_token", + AsyncMock(side_effect=asyncio.CancelledError()), + ), + ): + with pytest.raises(asyncio.CancelledError): + await refund_sweep_once() + + refund = (await _load(session_factory))["cancelled"] + assert refund.swept is False + assert refund.sweep_started_at == 1000 + + +@pytest.mark.asyncio +async def test_checkpoint_failure_retains_claim_and_stale_retry_records_sweep( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _insert( + session_factory, + CashuTransaction( + token="checkpoint-failure", + amount=1, + unit="sat", + type="out", + created_at=800, + ), + ) + real_set_state = wallet._set_refund_sweep_state + + async def fail_swept_checkpoint( + refund_id: str, + *, + predicates: tuple[object, ...] = (), + **values: object, + ) -> int: + if values.get("swept") is True: + raise RuntimeError("checkpoint unavailable") + return await real_set_state(refund_id, predicates=predicates, **values) + + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200), + patch("routstr.wallet.time.time", return_value=1000), + patch( + "routstr.wallet.recieve_token", AsyncMock(return_value=(1, "sat", "mint")) + ), + patch( + "routstr.wallet._set_refund_sweep_state", + side_effect=fail_swept_checkpoint, + ), + patch("routstr.wallet.logger.critical") as critical, + ): + await refund_sweep_once() + + retained = (await _load(session_factory))["checkpoint-failure"] + assert retained.swept is False + assert retained.collected is False + assert retained.sweep_started_at == 1000 + critical.assert_called_once() + + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200), + patch("routstr.wallet.time.time", return_value=1300), + patch( + "routstr.wallet.recieve_token", + AsyncMock(side_effect=RuntimeError("token already spent")), + ), + ): + await refund_sweep_once() + + recovered = (await _load(session_factory))["checkpoint-failure"] + assert recovered.swept is True + assert recovered.collected is False + assert recovered.sweep_started_at is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("redemption_succeeds", [True, False]) +async def test_expired_worker_cannot_overwrite_or_release_newer_claim( + session_factory: async_sessionmaker[AsyncSession], + redemption_succeeds: bool, +) -> None: + await _insert( + session_factory, + CashuTransaction( + token="reclaimed", + amount=1, + unit="sat", + type="out", + created_at=800, + ), + ) + + async def replace_claim(_token: str) -> tuple[int, str, str]: + async with session_factory() as session: + result = await session.exec( + select(CashuTransaction).where(CashuTransaction.token == "reclaimed") + ) + transaction = result.one() + transaction.sweep_started_at = 1100 + session.add(transaction) + await session.commit() + if not redemption_succeeds: + raise RuntimeError("mint unavailable") + return (1, "sat", "mint") + + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.time.time", return_value=1000), + patch("routstr.wallet.recieve_token", AsyncMock(side_effect=replace_claim)), + ): + await refund_sweep_once() + + reclaimed = (await _load(session_factory))["reclaimed"] + assert reclaimed.swept is False + assert reclaimed.collected is False + assert reclaimed.sweep_started_at == 1100 + + +@pytest.mark.asyncio +async def test_refund_sweep_recovers_stale_claim_without_misreporting_collection( + session_factory: async_sessionmaker[AsyncSession], +) -> None: + await _insert( + session_factory, + CashuTransaction( + token="stale", + amount=1, + unit="sat", + type="out", + created_at=800, + sweep_started_at=100, + ), + CashuTransaction( + token="active", + amount=1, + unit="sat", + type="out", + created_at=800, + sweep_started_at=950, + ), + ) + receive = AsyncMock(side_effect=RuntimeError("token already spent")) + with ( + patch("routstr.wallet.db.create_session", side_effect=session_factory), + patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100), + patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200), + patch("routstr.wallet.time.time", return_value=1000), + patch("routstr.wallet.recieve_token", receive), + ): + await refund_sweep_once() + + receive.assert_awaited_once_with("stale") + loaded = await _load(session_factory) + assert loaded["stale"].swept is True + assert loaded["stale"].collected is False + assert loaded["stale"].sweep_started_at is None + assert loaded["active"].swept is False + assert loaded["active"].sweep_started_at == 950 diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index a553b6d4..98a9e24f 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -62,6 +62,101 @@ def test_payout_settings_have_sensible_defaults() -> None: assert s.payout_interval_seconds == 900 +def test_database_pool_defaults_match_sqlalchemy_capacity() -> 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_recycle == 1800 + assert s.database_pool_pre_ping is False + assert s.database_pool_hold_warn_seconds == 10.0 + + +@pytest.mark.parametrize( + ("field", "bad_value"), + [ + ("database_pool_size", 0), + ("database_max_overflow", -1), + ("database_pool_timeout", 0), + ("database_pool_recycle", -1), + ("database_pool_hold_warn_seconds", 0), + ], +) +def test_database_pool_settings_reject_invalid_values( + field: str, bad_value: int +) -> None: + with pytest.raises(ValidationError): + Settings.parse_obj({field: bad_value}) + + +@pytest.mark.asyncio +async def test_database_pool_fields_are_env_only_not_persisted( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """DB pool sizing is infrastructure the node needs *before* it can read the + DB, so it can never be configured from the DB — it must never be written to + the settings blob, and a stale/injected DB value must never shadow env. + """ + monkeypatch.setenv("DATABASE_POOL_SIZE", "7") + + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with AsyncSession(engine, expire_on_commit=False) as session: + s = await SettingsService.initialize(session) + + # The env value is live for runtime consumers... + assert s.database_pool_size == 7 + # ...but pool sizing is never written to the settings blob. + blob = await _read_settings_blob(session) + for field in ( + "database_pool_size", + "database_max_overflow", + "database_pool_timeout", + "database_pool_recycle", + "database_pool_pre_ping", + "database_pool_hold_warn_seconds", + ): + assert field not in blob + + # Even a stale blob that somehow carries a pool value must not win: env + # stays authoritative on the next initialize. + await session.exec( # type: ignore + text("UPDATE settings SET data = :d WHERE id = 1").bindparams( + d=json.dumps({"database_pool_size": 99}) + ) + ) + await session.commit() + again = await SettingsService.initialize(session) + assert again.database_pool_size == 7 + + +@pytest.mark.asyncio +async def test_update_does_not_apply_env_only_fields_to_live_settings( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """DB pool sizing is env-only: a settings update must neither persist it nor + mutate the live value. The engine pool is already built at boot from env, so + a UI/API update carrying a pool value must not make the live setting diverge + from the running pool. + """ + monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False) + monkeypatch.setattr(settings, "database_pool_size", 5) + + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with AsyncSession(engine, expire_on_commit=False) as session: + await SettingsService.initialize(session) + await SettingsService.update( + {"database_pool_size": 99, "name": "PoolTweaker"}, session + ) + + # 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 + # ...and it is never written to the settings blob. + blob = await _read_settings_blob(session) + assert "database_pool_size" not in blob + + @pytest.mark.parametrize( "field,bad_value", [