From ff55788e2d4a4b297239119f5264dd1c79cbec26 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 26 Jul 2026 12:44:37 +0200 Subject: [PATCH] fix: address PR 634 review feedback --- .env.example | 12 +- routstr/core/admin.py | 10 +- routstr/core/db.py | 97 ++++--- routstr/core/settings.py | 15 +- routstr/lightning.py | 119 +++++---- routstr/wallet.py | 240 ++++++++++++------ .../test_lightning_invoice_constraints.py | 101 +++++++- tests/unit/test_admin_transactions.py | 35 +++ tests/unit/test_balances_by_mint_and_unit.py | 17 +- tests/unit/test_db_pool_config.py | 48 +++- tests/unit/test_fee_payout_crash_safety.py | 50 ++++ tests/unit/test_periodic_payout.py | 17 +- tests/unit/test_refund_sweep.py | 111 ++++++++ tests/unit/test_settings.py | 8 +- 14 files changed, 664 insertions(+), 216 deletions(-) create mode 100644 tests/unit/test_admin_transactions.py diff --git a/.env.example b/.env.example index ba556919..9f0c0bbc 100644 --- a/.env.example +++ b/.env.example @@ -23,14 +23,14 @@ 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. Defaults intentionally bound capacity and fail quickly so -# leaked connections are visible instead of hidden by overflow connections. -# Keep total capacity across all workers below the database connection limit. +# 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=0 -# DATABASE_POOL_TIMEOUT=5 +# DATABASE_MAX_OVERFLOW=10 +# DATABASE_POOL_TIMEOUT=30 # DATABASE_POOL_RECYCLE=1800 -# DATABASE_POOL_PRE_PING=true +# 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 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 0ee04933..9e0c5973 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -22,6 +22,7 @@ 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__) @@ -29,18 +30,13 @@ DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: - """Build an async engine from validated, environment-only pool settings.""" - from .settings import settings - + """Build and instrument an async engine from environment-only settings.""" url = make_url(database_url) - is_memory_sqlite = url.get_backend_name() == "sqlite" and url.database in { - None, - "", - ":memory:", - } - options: dict[str, int | float | bool] = { - "pool_pre_ping": settings.database_pool_pre_ping, - } + 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, @@ -52,44 +48,45 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: logger.info( "Database pool configured", extra={ - "database_url_backend": url.get_backend_name(), + "database_url_backend": backend, "in_memory_sqlite": is_memory_sqlite, **options, }, ) - return create_async_engine(database_url, echo=False, **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() -from .settings import settings as _runtime_settings # noqa: E402 - -_POOL_HOLD_WARN_SECONDS = _runtime_settings.database_pool_hold_warn_seconds - - -@event.listens_for(engine.sync_engine, "checkout") -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] - - -@event.listens_for(engine.sync_engine, "checkin") -def _record_pool_checkin(dbapi_connection: object, connection_record: object) -> None: - checked_out_at = connection_record.info.pop("routstr_checked_out_at", None) # type: ignore[attr-defined] - if checked_out_at is None: - return - held_seconds = time.monotonic() - checked_out_at - if held_seconds >= _POOL_HOLD_WARN_SECONDS: - logger.warning( - "Database connection held longer than threshold", - extra={ - "held_seconds": round(held_seconds, 3), - "threshold_seconds": _POOL_HOLD_WARN_SECONDS, - "pool_status": engine.pool.status(), - }, - ) - class ApiKey(SQLModel, table=True): # type: ignore __tablename__ = "api_keys" @@ -284,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 @@ -783,6 +783,19 @@ async def complete_routstr_fee_payout( return result.rowcount == 1 +async def balance_for_mint_and_unit( + db_session: AsyncSession, mint_url: str, unit: str +) -> int: + """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]: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 34c6d9c1..0caffd16 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -106,14 +106,17 @@ class Settings(BaseSettings): default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS" ) - # Database connection-pool controls (advanced). Keep the default pool - # bounded and fail quickly: increasing it can hide leaked connections and - # makes SQLite lock contention worse. These fields are env-only below. + # 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=0, ge=0, env="DATABASE_MAX_OVERFLOW") - database_pool_timeout: float = Field(default=5.0, gt=0, env="DATABASE_POOL_TIMEOUT") + 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=True, env="DATABASE_POOL_PRE_PING") + 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" ) diff --git a/routstr/lightning.py b/routstr/lightning.py index 5bba7593..554925f2 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -225,48 +225,84 @@ async def check_invoice_payment( 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() + # 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 any mint I/O 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: + logger.error( + "Topup invoice target API key was not found; skipping mint", + extra={"invoice_id": invoice_id}, + ) + return + wallet = await get_wallet(settings.primary_mint, "sat") - mint_status = await wallet.get_mint_quote(invoice.payment_hash) + mint_status = await wallet.get_mint_quote(invoice_payment_hash) if not mint_status.paid: return # 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) + await wallet.mint(invoice_amount_sats, quote_id=invoice_payment_hash) minted = True - if invoice_purpose == "create": - api_key = await _create_api_key_record(invoice, session) - invoice.api_key_hash = api_key.hashed_key - elif invoice_purpose == "topup": - await _credit_topup_record(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) - # Conditional transition guards against double-credit: the credit - # above and this status flip commit atomically, and a lost race - # rolls both back. - paid_at = int(time.time()) - finalized = await session.exec( # type: ignore[call-overload] - update(LightningInvoice) - .where( - col(LightningInvoice.id) == invoice_id, - col(LightningInvoice.status) == "pending", + # 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, + ) ) - .values( - status="paid", - paid_at=paid_at, - api_key_hash=invoice.api_key_hash, - ) - ) - if finalized.rowcount != 1: - await session.rollback() - await session.refresh(invoice) - return - await session.commit() + 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. + invoice.api_key_hash = committed_invoice.api_key_hash + invoice.status = committed_invoice.status + invoice.paid_at = committed_invoice.paid_at + return + await finalization_session.commit() + + # Only publish finalized values to the caller-owned object after the + # owned transaction has committed successfully. + invoice.api_key_hash = finalized_api_key_hash invoice.status = "paid" invoice.paid_at = paid_at @@ -274,20 +310,17 @@ async def check_invoice_payment( "Lightning invoice paid", extra={ "invoice_id": invoice_id, - "amount_sats": invoice.amount_sats, + "amount_sats": invoice_amount_sats, "purpose": invoice_purpose, - "api_key_hash": invoice.api_key_hash[:8] + "..." - if invoice.api_key_hash + "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. - try: - await asyncio.shield(session.rollback()) - except Exception: - logger.warning("Rollback failed during invoice check cleanup") + # 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", @@ -331,22 +364,6 @@ async def _credit_topup_record( raise ValueError("Associated API key not found") -async def create_api_key_from_invoice( - invoice: LightningInvoice, session: AsyncSession -) -> ApiKey: - wallet = await get_wallet(settings.primary_mint, "sat") - await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) - return await _create_api_key_record(invoice, session) - - -async def topup_api_key_from_invoice( - invoice: LightningInvoice, session: AsyncSession -) -> None: - wallet = await get_wallet(settings.primary_mint, "sat") - await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash) - await _credit_topup_record(invoice, session) - - INVOICE_WATCH_INTERVAL_SECONDS = 5 INVOICE_WATCH_BATCH_LIMIT = 100 diff --git a/routstr/wallet.py b/routstr/wallet.py index ac2fda3d..e2a20a7c 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} ) @@ -838,9 +857,7 @@ async def fetch_all_balances( # 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: 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() user_balances: dict[tuple[str, str], int] = {} liabilities_error: str | None = None @@ -865,7 +882,7 @@ async def fetch_all_balances( 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 = { @@ -900,13 +917,13 @@ async def fetch_all_balances( total_wallet_balance_sats += ( detail["wallet_balance"] if unit == "sat" - else detail["wallet_balance"] // 1000 + 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"]) ) if liabilities_error is None: @@ -938,9 +955,7 @@ 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() for mint_url in mint_urls: for unit in ["sat", "msat"]: @@ -953,24 +968,47 @@ async def periodic_payout() -> None: ) proofs = await slow_filter_spend_proofs(proofs, wallet) await asyncio.sleep(5) - # 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. + 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: - balances = await db.balances_by_mint_and_unit( - session, [mint_url], [unit] + user_balance = await db.balance_for_mint_and_unit( + session, mint_url, unit ) - user_balance = balances.get((mint_url, unit), 0) + 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 = user_balance // 1000 + 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 settings.min_payout_sat * 1000 + else _sats_to_msats(settings.min_payout_sat) ) if available_balance > min_amount: amount_received = await raw_send_to_lnurl( @@ -1005,20 +1043,24 @@ async def periodic_payout() -> None: ) -async def _set_refund_sweep_state(refund_id: str, **values: object) -> None: +async def _set_refund_sweep_state( + refund_id: str, + *, + predicates: tuple[typing.Any, ...] = (), + **values: object, +) -> int: async with db.create_session() as session: - await session.exec( # type: ignore[call-overload] + result = await session.exec( # type: ignore[call-overload] update(db.CashuTransaction) - .where(col(db.CashuTransaction.id) == refund_id) + .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_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 ) @@ -1036,69 +1078,109 @@ async def _refund_sweep_once(cutoff: int) -> None: for refund in refunds: reclaimed_stale_claim = refund.sweep_started_at is not None claim_started_at = int(time.time()) - async with db.create_session() as session: - claim = await session.exec( # type: ignore[call-overload] - update(db.CashuTransaction) - .where( - col(db.CashuTransaction.id) == refund.id, - col(db.CashuTransaction.swept) == False, # noqa: E712 - col(db.CashuTransaction.collected) == False, # noqa: E712 - claim_available, - ) - .values(sweep_started_at=claim_started_at) - ) - await session.commit() - if claim.rowcount != 1: + 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) - await _set_refund_sweep_state( - refund.id, swept=True, sweep_started_at=None - ) - logger.info( - "Swept uncollected refund", - extra={ - "id": refund.id, - "amount": refund.amount, - "unit": refund.unit, - }, + 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={ + "id": refund.id, + "amount": refund.amount, + "unit": refund.unit, + }, + ) + else: + logger.critical( + "Refund token swept after claim ownership changed; manual reconciliation required", + extra={"id": refund.id}, + ) except BaseException as e: + if redeemed: + # The token is already in the node wallet. Retain the claim so + # a stale retry classifies "already spent" as swept, never as a + # client collection. + logger.critical( + "Refund token swept but 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. - await _set_refund_sweep_state( - refund.id, swept=True, sweep_started_at=None + updated = await _set_refund_sweep_state( + refund.id, + predicates=(claim_owned,), + swept=True, + sweep_started_at=None, ) else: - await _set_refund_sweep_state( + updated = await _set_refund_sweep_state( refund.id, + predicates=(claim_owned,), collected=True, swept=False, sweep_started_at=None, ) - logger.info( - "Refund token was already spent", - extra={ - "id": refund.id, - "reclaimed_stale_claim": reclaimed_stale_claim, - }, - ) + if updated == 1: + logger.info( + "Refund token was already spent", + extra={ + "id": refund.id, + "reclaimed_stale_claim": reclaimed_stale_claim, + }, + ) + else: + logger.warning( + "Refund claim ownership changed before spent-token checkpoint", + extra={"id": refund.id}, + ) else: - # Cancellation is a known-no-send outcome at this boundary, so - # release the lease under shield and let a later sweep retry. - await asyncio.shield( - _set_refund_sweep_state(refund.id, sweep_started_at=None) + # Cancellation or a transient pre-redemption failure is a + # known-no-send outcome, so release the lease for a later retry. + released = await asyncio.shield( + _set_refund_sweep_state( + refund.id, + predicates=(claim_owned,), + sweep_started_at=None, + ) ) if not isinstance(e, Exception): raise logger.warning( "Failed to sweep refund", - extra={"id": refund.id, "error": str(e)}, + extra={ + "id": refund.id, + "error": str(e), + "claim_released": released == 1, + }, ) @@ -1145,10 +1227,10 @@ async def periodic_routstr_fee_payout() -> None: ) continue - accumulated_sats = fee.accumulated_msats // 1000 + accumulated_sats = _msats_to_sats(fee.accumulated_msats) if accumulated_sats < ROUTSTR_FEE_DEFAULT_PAYOUT: continue - paid_msats = accumulated_sats * 1000 + paid_msats = _sats_to_msats(accumulated_sats) # Wallet/proof preparation cannot send funds, so do it before the # durable checkpoint. A preparation failure must not strand an @@ -1180,10 +1262,20 @@ async def periodic_routstr_fee_payout() -> None: ) continue - async with db.create_session() as session: - payout_completed = await db.complete_routstr_fee_payout( - session, paid_msats + try: + async with db.create_session() as session: + payout_completed = await db.complete_routstr_fee_payout( + session, paid_msats + ) + except BaseException as e: + logger.critical( + "Routstr fee payout sent but checkpoint was not completed", + extra={"payout_in_progress_msats": paid_msats}, + exc_info=isinstance(e, Exception), ) + if not isinstance(e, Exception): + raise + continue if not payout_completed: logger.critical( "Routstr fee payout sent but checkpoint was not completed", diff --git a/tests/integration/test_lightning_invoice_constraints.py b/tests/integration/test_lightning_invoice_constraints.py index a7fa1384..b4e12a7f 100644 --- a/tests/integration/test_lightning_invoice_constraints.py +++ b/tests/integration/test_lightning_invoice_constraints.py @@ -3,8 +3,8 @@ 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 @@ -14,11 +14,12 @@ import time 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: @@ -104,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) @@ -120,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) @@ -137,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) @@ -148,6 +149,7 @@ async def test_created_key_receives_validity_date( @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: @@ -173,6 +175,7 @@ async def test_payment_check_releases_connection_during_mint_quote( @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: @@ -227,6 +230,7 @@ async def test_concurrent_payment_checks_mint_and_credit_invoice_once( @pytest.mark.asyncio async def test_failed_mint_keeps_invoice_pending_for_retry( 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: @@ -251,20 +255,62 @@ async def test_failed_mint_keeps_invoice_pending_for_retry( @pytest.mark.asyncio -async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation( +async def test_missing_topup_target_is_rejected_before_mint( integration_engine: AsyncEngine, + patched_db_engine: None, ) -> None: - invoice = _make_invoice(id="inv_finalize_failure", status="pending", paid_at=None) + invoice = _make_invoice( + id="inv_missing_topup_target", + status="pending", + paid_at=None, + purpose="topup", + api_key_hash="pruned-key", + ) async with AsyncSession(integration_engine, expire_on_commit=False) as setup: setup.add(invoice) await setup.commit() + get_wallet = 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", get_wallet): + from routstr.lightning import check_invoice_payment + + await check_invoice_payment(stored, session) + + get_wallet.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 == "pending" + + +@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( @@ -276,6 +322,15 @@ async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation( 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) @@ -291,7 +346,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) @@ -304,6 +359,7 @@ async def test_created_key_without_constraints_has_none_fields( @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.""" @@ -315,9 +371,15 @@ async def test_db_guard_credits_once_when_both_mints_succeed( 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(key) - setup.add(invoice) + setup.add_all([key, invoice, sibling]) await setup.commit() wallet = MagicMock() @@ -334,8 +396,10 @@ async def test_db_guard_credits_once_when_both_mints_succeed( 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)): @@ -346,6 +410,21 @@ async def test_db_guard_credits_once_when_both_mints_succeed( 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 wallet.mint.await_count == 2 async with AsyncSession(integration_engine, expire_on_commit=False) as verify: stored_invoice = await verify.get(LightningInvoice, invoice.id) 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 index bb1b3c81..248c43a8 100644 --- a/tests/unit/test_balances_by_mint_and_unit.py +++ b/tests/unit/test_balances_by_mint_and_unit.py @@ -13,7 +13,11 @@ from sqlalchemy.pool import StaticPool from sqlmodel import SQLModel from sqlmodel.ext.asyncio.session import AsyncSession -from routstr.core.db import ApiKey, balances_by_mint_and_unit +from routstr.core.db import ( + ApiKey, + balance_for_mint_and_unit, + balances_by_mint_and_unit, +) def _make_engine() -> AsyncEngine: @@ -92,6 +96,17 @@ async def test_excludes_rows_with_null_mint_or_currency(session: AsyncSession) - 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") diff --git a/tests/unit/test_db_pool_config.py b/tests/unit/test_db_pool_config.py index a2838621..8f5d0367 100644 --- a/tests/unit/test_db_pool_config.py +++ b/tests/unit/test_db_pool_config.py @@ -1,4 +1,3 @@ -from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest @@ -43,17 +42,44 @@ async def test_memory_sqlite_keeps_static_pool( await engine.dispose() -def test_pool_observer_warns_when_connection_is_held_too_long( +def test_non_sqlite_backend_enables_pre_ping_automatically( monkeypatch: pytest.MonkeyPatch, ) -> None: - record = SimpleNamespace(info={}) - monotonic = MagicMock(side_effect=[100.0, 112.5]) - monkeypatch.setattr(db.time, "monotonic", monotonic) - monkeypatch.setattr(db, "_POOL_HOLD_WARN_SECONDS", 10.0) + monkeypatch.setattr(settings, "database_pool_pre_ping", False) + fake_engine = MagicMock() - with patch.object(db.logger, "warning") as warning: - db._record_pool_checkout(None, record, None) - db._record_pool_checkin(None, record) + 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") - warning.assert_called_once() - assert warning.call_args.kwargs["extra"]["held_seconds"] == 12.5 + 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 7a0cd94a..0768db16 100644 --- a/tests/unit/test_fee_payout_crash_safety.py +++ b/tests/unit/test_fee_payout_crash_safety.py @@ -252,6 +252,56 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non critical.assert_called_once() +@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 diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index 69d13ee8..953f06cb 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -77,8 +77,8 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch( - "routstr.wallet.db.balances_by_mint_and_unit", - AsyncMock(return_value={}), + "routstr.wallet.db.balance_for_mint_and_unit", + AsyncMock(return_value=0), ), patch("routstr.wallet.raw_send_to_lnurl", raw_send), ): @@ -131,8 +131,8 @@ async def test_periodic_payout_releases_session_before_slow_mint_send() -> None: AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch( - "routstr.wallet.db.balances_by_mint_and_unit", - AsyncMock(return_value={}), + "routstr.wallet.db.balance_for_mint_and_unit", + AsyncMock(return_value=0), ), patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)), ): @@ -173,8 +173,8 @@ async def test_periodic_payout_isolates_failing_mint() -> None: AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch( - "routstr.wallet.db.balances_by_mint_and_unit", - AsyncMock(return_value={}), + "routstr.wallet.db.balance_for_mint_and_unit", + AsyncMock(return_value=0), ), patch("routstr.wallet.raw_send_to_lnurl", raw_send), ): @@ -220,10 +220,11 @@ async def test_periodic_payout_handles_session_creation_failure() -> None: await periodic_payout() # The liability session is opened per mint/unit (sat + msat), and each - # failure is isolated to its own iteration rather than aborting the cycle. + # 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 sending payout: RuntimeError" + assert message == "Error in periodic payout cycle: RuntimeError" assert extra["error"] == "db unavailable" diff --git a/tests/unit/test_refund_sweep.py b/tests/unit/test_refund_sweep.py index 3a8a68e2..b43c3152 100644 --- a/tests/unit/test_refund_sweep.py +++ b/tests/unit/test_refund_sweep.py @@ -9,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 @@ -185,6 +186,116 @@ async def test_refund_sweep_releases_claim_on_cancellation( assert refund.sweep_started_at is None +@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], diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index 965be8e6..98a9e24f 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -62,13 +62,13 @@ def test_payout_settings_have_sensible_defaults() -> None: assert s.payout_interval_seconds == 900 -def test_database_pool_defaults_are_bounded_and_fail_fast() -> None: +def test_database_pool_defaults_match_sqlalchemy_capacity() -> None: s = Settings() assert s.database_pool_size == 5 - assert s.database_max_overflow == 0 - assert s.database_pool_timeout == 5.0 + 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 True + assert s.database_pool_pre_ping is False assert s.database_pool_hold_warn_seconds == 10.0