mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-30 15:26:14 +00:00
fix: address PR 634 review feedback
This commit is contained in:
+6
-6
@@ -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
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
+55
-42
@@ -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]:
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
+68
-51
@@ -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
|
||||
|
||||
|
||||
+166
-74
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user