fix: address PR 634 review feedback

This commit is contained in:
9qeklajc
2026-07-26 12:44:37 +02:00
parent 1131c2d583
commit ff55788e2d
14 changed files with 664 additions and 216 deletions
+6 -6
View File
@@ -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
+8 -2
View File
@@ -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,
}
+44 -31
View File
@@ -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
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(
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]
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 >= _POOL_HOLD_WARN_SECONDS:
if held_seconds >= 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(),
"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
__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]:
+9 -6
View File
@@ -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"
)
+52 -35
View File
@@ -225,32 +225,54 @@ 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
# 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, session)
invoice.api_key_hash = api_key.hashed_key
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, session)
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.
# rolls both back in the owned finalization session.
paid_at = int(time.time())
finalized = await session.exec( # type: ignore[call-overload]
finalized = await finalization_session.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(
col(LightningInvoice.id) == invoice_id,
@@ -259,14 +281,28 @@ async def check_invoice_payment(
.values(
status="paid",
paid_at=paid_at,
api_key_hash=invoice.api_key_hash,
api_key_hash=finalized_api_key_hash,
)
)
if finalized.rowcount != 1:
await session.rollback()
await session.refresh(invoice)
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 session.commit()
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
+145 -53
View File
@@ -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.
async with db.create_session() as session:
balances = await db.balances_by_mint_and_unit(
session, [mint_url], [unit]
except Exception as e:
logger.error(
f"Error sending payout: {type(e).__name__}",
extra={
"error": str(e),
"mint_url": mint_url,
"unit": unit,
},
)
user_balance = balances.get((mint_url, unit), 0)
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
)
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,26 +1078,30 @@ 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,
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,
)
.values(sweep_started_at=claim_started_at)
)
await session.commit()
if claim.rowcount != 1:
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
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={
@@ -1064,23 +1110,46 @@ async def _refund_sweep_once(cutoff: int) -> None:
"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,
)
if updated == 1:
logger.info(
"Refund token was already spent",
extra={
@@ -1089,16 +1158,29 @@ async def _refund_sweep_once(cutoff: int) -> None:
},
)
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)
logger.warning(
"Refund claim ownership changed before spent-token checkpoint",
extra={"id": refund.id},
)
else:
# 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
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",
@@ -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)
+35
View File
@@ -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]
+16 -1
View File
@@ -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")
+36 -10
View File
@@ -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, "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:
db._record_pool_checkout(None, record, None)
db._record_pool_checkin(None, record)
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")
warning.assert_called_once()
assert warning.call_args.kwargs["extra"]["held_seconds"] == 12.5
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()
@@ -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
+9 -8
View File
@@ -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"
+111
View File
@@ -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],
+4 -4
View File
@@ -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