merge: combine PR #632 with broader pool-exhaustion fixes

This commit is contained in:
9qeklajc
2026-07-24 23:22:01 +02:00
12 changed files with 675 additions and 183 deletions
+8 -1
View File
@@ -22,13 +22,19 @@ ROUTSTR_SECRET_KEY=
# Database # Database
# DATABASE_URL=sqlite+aiosqlite:///keys.db # DATABASE_URL=sqlite+aiosqlite:///keys.db
# Keep total pool capacity across all workers below the database connection limit. # 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.
# DATABASE_POOL_SIZE=5 # DATABASE_POOL_SIZE=5
# DATABASE_MAX_OVERFLOW=0 # DATABASE_MAX_OVERFLOW=0
# DATABASE_POOL_TIMEOUT=5 # DATABASE_POOL_TIMEOUT=5
# DATABASE_POOL_RECYCLE=1800 # DATABASE_POOL_RECYCLE=1800
# DATABASE_POOL_PRE_PING=true # DATABASE_POOL_PRE_PING=true
# Warn when a checkout is held this many seconds.
# DATABASE_POOL_HOLD_WARN_SECONDS=10 # DATABASE_POOL_HOLD_WARN_SECONDS=10
# SQLite serialises writes; increasing its pool can trade pool timeouts for
# "database is locked" errors rather than increasing write throughput.
# Node Information # Node Information
# NAME=My Routstr Node # NAME=My Routstr Node
@@ -40,6 +46,7 @@ ROUTSTR_SECRET_KEY=
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me" # CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me"
# MINT_OPERATION_CONCURRENCY=4 # MINT_OPERATION_CONCURRENCY=4
# RECEIVE_LN_ADDRESS= # RECEIVE_LN_ADDRESS=
# REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS=900
# Custom Pricing Configuration # Custom Pricing Configuration
# MODEL_BASED_PRICING=true # MODEL_BASED_PRICING=true
@@ -0,0 +1,37 @@
"""add recoverable refund sweep claim lease
Revision ID: 8a1b2c3d4e5f
Revises: 7f2843d3f4e4
Create Date: 2026-07-24 23:15:00.000000
"""
import sqlalchemy as sa
from alembic import op
revision = "8a1b2c3d4e5f"
down_revision = "7f2843d3f4e4"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
columns = {
column["name"]
for column in sa.inspect(conn).get_columns("cashu_transactions")
}
if "sweep_started_at" not in columns:
op.add_column(
"cashu_transactions",
sa.Column("sweep_started_at", sa.Integer(), nullable=True),
)
def downgrade() -> None:
conn = op.get_bind()
columns = {
column["name"]
for column in sa.inspect(conn).get_columns("cashu_transactions")
}
if "sweep_started_at" in columns:
op.drop_column("cashu_transactions", "sweep_started_at")
+36 -54
View File
@@ -15,6 +15,7 @@ from alembic.util.exc import CommandError
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_ from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_
from sqlalchemy.engine import make_url from sqlalchemy.engine import make_url
from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.ext.asyncio import AsyncEngine
from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlalchemy.orm import aliased from sqlalchemy.orm import aliased
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
@@ -27,58 +28,43 @@ logger = get_logger(__name__)
DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
def _env_bool(name: str, default: bool) -> bool: def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
raw = os.environ.get(name) """Build an async engine from validated, environment-only pool settings."""
if raw is None: from .settings import settings
return default
normalized = raw.strip().lower()
if normalized in {"1", "true", "yes", "on"}:
return True
if normalized in {"0", "false", "no", "off"}:
return False
raise ValueError(f"{name} must be a boolean")
def _engine_options(database_url: str) -> dict[str, int | float | bool]:
"""Build a bounded pool configuration, preserving SQLite memory semantics."""
options: dict[str, int | float | bool] = {
"pool_pre_ping": _env_bool("DATABASE_POOL_PRE_PING", True)
}
url = make_url(database_url) url = make_url(database_url)
is_memory_sqlite = url.get_backend_name() == "sqlite" and url.database in { is_memory_sqlite = url.get_backend_name() == "sqlite" and url.database in {
None, None,
"", "",
":memory:", ":memory:",
} }
if is_memory_sqlite: options: dict[str, int | float | bool] = {
return options "pool_pre_ping": settings.database_pool_pre_ping,
}
if not is_memory_sqlite:
options.update(
pool_size=settings.database_pool_size,
max_overflow=settings.database_max_overflow,
pool_timeout=settings.database_pool_timeout,
pool_recycle=settings.database_pool_recycle,
)
pool_size = int(os.environ.get("DATABASE_POOL_SIZE", "5")) logger.info(
max_overflow = int(os.environ.get("DATABASE_MAX_OVERFLOW", "0")) "Database pool configured",
pool_timeout = float(os.environ.get("DATABASE_POOL_TIMEOUT", "5")) extra={
pool_recycle = int(os.environ.get("DATABASE_POOL_RECYCLE", "1800")) "database_url_backend": url.get_backend_name(),
if pool_size < 1: "in_memory_sqlite": is_memory_sqlite,
raise ValueError("DATABASE_POOL_SIZE must be at least 1") **options,
if max_overflow < 0: },
raise ValueError("DATABASE_MAX_OVERFLOW cannot be negative")
if pool_timeout <= 0:
raise ValueError("DATABASE_POOL_TIMEOUT must be positive")
if pool_recycle < 0:
raise ValueError("DATABASE_POOL_RECYCLE cannot be negative")
options.update(
pool_size=pool_size,
max_overflow=max_overflow,
pool_timeout=pool_timeout,
pool_recycle=pool_recycle,
) )
return options return create_async_engine(database_url, echo=False, **options)
engine = create_async_engine(DATABASE_URL, echo=False, **_engine_options(DATABASE_URL)) engine = create_db_engine()
_POOL_HOLD_WARN_SECONDS = float(os.environ.get("DATABASE_POOL_HOLD_WARN_SECONDS", "10")) from .settings import settings as _runtime_settings # noqa: E402
if _POOL_HOLD_WARN_SECONDS <= 0:
raise ValueError("DATABASE_POOL_HOLD_WARN_SECONDS must be positive") _POOL_HOLD_WARN_SECONDS = _runtime_settings.database_pool_hold_warn_seconds
@event.listens_for(engine.sync_engine, "checkout") @event.listens_for(engine.sync_engine, "checkout")
@@ -441,6 +427,10 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
) )
collected: bool = Field(default=False) collected: bool = Field(default=False)
swept: bool = Field(default=False) swept: bool = Field(default=False)
sweep_started_at: int | None = Field(
default=None,
description="Unix timestamp for a recoverable refund-sweep claim",
)
source: str = Field( source: str = Field(
default="x-cashu", default="x-cashu",
description="Payment source: x-cashu or apikey", description="Payment source: x-cashu or apikey",
@@ -793,20 +783,12 @@ async def complete_routstr_fee_payout(
return result.rowcount == 1 return result.rowcount == 1
async def balances_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str
) -> int:
query = select(func.sum(ApiKey.balance)).where(
ApiKey.refund_mint_url == mint_url, ApiKey.refund_currency == unit
)
result = await db_session.exec(query)
return result.one() or 0
async def balances_by_mint_and_unit( async def balances_by_mint_and_unit(
db_session: AsyncSession, db_session: AsyncSession, mint_urls: list[str], units: list[str]
) -> dict[tuple[str, str], int]: ) -> dict[tuple[str, str], int]:
"""Return all user liabilities in one query, grouped by mint and unit.""" """Return requested user liabilities grouped by mint and unit."""
if not mint_urls or not units:
return {}
query = ( query = (
select( select(
col(ApiKey.refund_mint_url), col(ApiKey.refund_mint_url),
@@ -814,8 +796,8 @@ async def balances_by_mint_and_unit(
func.sum(ApiKey.balance), func.sum(ApiKey.balance),
) )
.where( .where(
col(ApiKey.refund_mint_url).is_not(None), col(ApiKey.refund_mint_url).in_(mint_urls),
col(ApiKey.refund_currency).is_not(None), col(ApiKey.refund_currency).in_(units),
) )
.group_by(col(ApiKey.refund_mint_url), col(ApiKey.refund_currency)) .group_by(col(ApiKey.refund_mint_url), col(ApiKey.refund_currency))
) )
+52 -4
View File
@@ -102,6 +102,21 @@ class Settings(BaseSettings):
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH") enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS") refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
refund_sweep_claim_timeout_seconds: int = Field(
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS"
)
# Database connection-pool controls (advanced). 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_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_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_hold_warn_seconds: float = Field(
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
)
# Logging # Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL") log_level: str = Field(default="INFO", env="LOG_LEVEL")
@@ -147,10 +162,32 @@ def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
# ``routstr.core.vault``. # ``routstr.core.vault``.
SECRET_FIELDS = frozenset({"admin_password", "nsec"}) SECRET_FIELDS = frozenset({"admin_password", "nsec"})
# Infrastructure the node needs *before* it can open a DB session — so it can
# never be configured from the DB (chicken-and-egg) and stays env-only. Unlike
# secrets (owned by bootstrap), these are excluded so the DB settings blob can
# neither store nor shadow them; env is always authoritative.
ENV_ONLY_FIELDS = frozenset(
{
"database_pool_size",
"database_max_overflow",
"database_pool_timeout",
"database_pool_recycle",
"database_pool_pre_ping",
"database_pool_hold_warn_seconds",
}
)
_NON_PERSISTED_FIELDS = SECRET_FIELDS | ENV_ONLY_FIELDS
def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]: def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]:
"""Return a copy of ``data`` without any secret fields (for persistence).""" """Return a copy of ``data`` without secret or env-only fields.
return {k: v for k, v in data.items() if k not in SECRET_FIELDS}
Both are kept out of the persisted settings blob: secrets for confidentiality,
env-only fields (e.g. DB pool sizing) because they must never be sourced from
the database.
"""
return {k: v for k, v in data.items() if k not in _NON_PERSISTED_FIELDS}
def _apply_to_live_settings(data: dict[str, Any]) -> None: def _apply_to_live_settings(data: dict[str, Any]) -> None:
@@ -330,7 +367,13 @@ class SettingsService:
valid_fields = set(env_resolved.dict().keys()) valid_fields = set(env_resolved.dict().keys())
merged_dict: dict[str, Any] = dict(env_resolved.dict()) merged_dict: dict[str, Any] = dict(env_resolved.dict())
merged_dict.update( merged_dict.update(
{k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields} {
k: v
for k, v in db_json.items()
if v not in (None, "", [], {})
and k in valid_fields
and k not in ENV_ONLY_FIELDS
}
) )
merged_dict = Settings(**merged_dict).dict() merged_dict = Settings(**merged_dict).dict()
@@ -397,8 +440,13 @@ class SettingsService:
) )
) )
await db_session.commit() await db_session.commit()
# Update in-place # Update in-place. Env-only fields (e.g. DB pool sizing) are never
# applied here: the engine pool is already built at boot from env,
# so letting an update mutate the live value would only make it
# diverge from the running pool.
for k, v in candidate.dict().items(): for k, v in candidate.dict().items():
if k in ENV_ONLY_FIELDS:
continue
setattr(settings, k, v) setattr(settings, k, v)
cls._current = settings cls._current = settings
return settings return settings
+102 -59
View File
@@ -836,8 +836,22 @@ async def fetch_all_balances(
if units is None: if units is None:
units = ["sat", "msat"] units = ["sat", "msat"]
async with db.create_session() as session: # Received tokens are stored against primary_mint even when cashu_mints is
user_balances = await db.balances_by_mint_and_unit(session) # 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)
user_balances: dict[tuple[str, str], int] = {}
liabilities_error: str | None = None
try:
async with db.create_session() as session:
user_balances = await db.balances_by_mint_and_unit(
session, mint_urls, units
)
except Exception as e:
logger.error("Error reading user balances", extra={"error": str(e)})
liabilities_error = str(e)
mint_check_limit = asyncio.Semaphore(settings.mint_operation_concurrency) mint_check_limit = asyncio.Semaphore(settings.mint_operation_concurrency)
@@ -874,40 +888,37 @@ async def fetch_all_balances(
} }
return error_result return error_result
# Build the set of mints to inspect. Received tokens are stored against
# ``primary_mint`` (which defaults to a real mint even when ``cashu_mints``
# is empty), so include it as a fallback — otherwise a node that accepts
# payments would still report empty balances when ``cashu_mints`` is unset.
mint_urls: list[str] = list(settings.cashu_mints)
if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint)
tasks = [fetch_balance(mint_url, unit) for mint_url in mint_urls for unit in units] tasks = [fetch_balance(mint_url, unit) for mint_url in mint_urls for unit in units]
balance_details = list(await asyncio.gather(*tasks)) balance_details = list(await asyncio.gather(*tasks))
# Calculate totals
total_wallet_balance_sats = 0 total_wallet_balance_sats = 0
total_user_balance_sats = 0 total_user_balance_sats = 0
for detail in balance_details: for detail in balance_details:
if not detail.get("error"): if detail.get("error"):
# Convert to sats for total calculation continue
unit = detail["unit"] unit = detail["unit"]
proofs_balance_sats = ( total_wallet_balance_sats += (
detail["wallet_balance"] detail["wallet_balance"]
if unit == "sat" if unit == "sat"
else detail["wallet_balance"] // 1000 else detail["wallet_balance"] // 1000
) )
user_balance_sats = ( if liabilities_error is None:
total_user_balance_sats += (
detail["user_balance"] detail["user_balance"]
if unit == "sat" if unit == "sat"
else detail["user_balance"] // 1000 else detail["user_balance"] // 1000
) )
total_wallet_balance_sats += proofs_balance_sats if liabilities_error is None:
total_user_balance_sats += user_balance_sats owner_balance = total_wallet_balance_sats - total_user_balance_sats
else:
owner_balance = total_wallet_balance_sats - total_user_balance_sats # Custody remains knowable when the DB read fails, but the user/owner
# split does not. Never report unknown liabilities as owner profit.
owner_balance = 0
for detail in balance_details:
detail["user_balance"] = 0
detail["owner_balance"] = 0
detail.setdefault("error", liabilities_error)
return ( return (
balance_details, balance_details,
@@ -947,9 +958,10 @@ async def periodic_payout() -> None:
# the liability, shrinking the payout — never sending # the liability, shrinking the payout — never sending
# customer-backed funds as profit. # customer-backed funds as profit.
async with db.create_session() as session: async with db.create_session() as session:
user_balance = await db.balances_for_mint_and_unit( balances = await db.balances_by_mint_and_unit(
session, mint_url, unit session, [mint_url], [unit]
) )
user_balance = balances.get((mint_url, unit), 0)
if unit == "sat": if unit == "sat":
user_balance = user_balance // 1000 user_balance = user_balance // 1000
proofs_balance = sum(proof.amount for proof in proofs) proofs_balance = sum(proof.amount for proof in proofs)
@@ -993,20 +1005,37 @@ async def periodic_payout() -> None:
) )
async def _set_refund_sweep_state(refund_id: str, **values: object) -> None:
async with db.create_session() as session:
await session.exec( # type: ignore[call-overload]
update(db.CashuTransaction)
.where(col(db.CashuTransaction.id) == refund_id)
.values(**values)
)
await session.commit()
async def _refund_sweep_once(cutoff: int) -> None: async def _refund_sweep_once(cutoff: int) -> None:
claim_cutoff = (
int(time.time()) - settings.refund_sweep_claim_timeout_seconds
)
claim_available = col(db.CashuTransaction.sweep_started_at).is_(None) | (
col(db.CashuTransaction.sweep_started_at) < claim_cutoff
)
async with db.create_session() as session: async with db.create_session() as session:
stmt = select(db.CashuTransaction).where( stmt = select(db.CashuTransaction).where(
db.CashuTransaction.type == "out", db.CashuTransaction.type == "out",
db.CashuTransaction.collected == False, # noqa: E712 db.CashuTransaction.collected == False, # noqa: E712
db.CashuTransaction.swept == False, # noqa: E712 db.CashuTransaction.swept == False, # noqa: E712
db.CashuTransaction.created_at < cutoff, db.CashuTransaction.created_at < cutoff,
claim_available,
) )
results = await session.exec(stmt) results = await session.exec(stmt)
refunds = results.all() refunds = results.all()
for refund in refunds: for refund in refunds:
# Claim the refund atomically before redeeming so a concurrent sweep reclaimed_stale_claim = refund.sweep_started_at is not None
# cannot redeem the same token and misreport it as client-collected. claim_started_at = int(time.time())
async with db.create_session() as session: async with db.create_session() as session:
claim = await session.exec( # type: ignore[call-overload] claim = await session.exec( # type: ignore[call-overload]
update(db.CashuTransaction) update(db.CashuTransaction)
@@ -1014,8 +1043,9 @@ async def _refund_sweep_once(cutoff: int) -> None:
col(db.CashuTransaction.id) == refund.id, col(db.CashuTransaction.id) == refund.id,
col(db.CashuTransaction.swept) == False, # noqa: E712 col(db.CashuTransaction.swept) == False, # noqa: E712
col(db.CashuTransaction.collected) == False, # noqa: E712 col(db.CashuTransaction.collected) == False, # noqa: E712
claim_available,
) )
.values(swept=True) .values(sweep_started_at=claim_started_at)
) )
await session.commit() await session.commit()
if claim.rowcount != 1: if claim.rowcount != 1:
@@ -1023,6 +1053,9 @@ async def _refund_sweep_once(cutoff: int) -> None:
try: try:
await recieve_token(refund.token) await recieve_token(refund.token)
await _set_refund_sweep_state(
refund.id, swept=True, sweep_started_at=None
)
logger.info( logger.info(
"Swept uncollected refund", "Swept uncollected refund",
extra={ extra={
@@ -1031,38 +1064,42 @@ async def _refund_sweep_once(cutoff: int) -> None:
"unit": refund.unit, "unit": refund.unit,
}, },
) )
except Exception as e: except BaseException as e:
error_msg = str(e).lower() error_msg = str(e).lower()
if "already spent" in error_msg: if isinstance(e, Exception) and "already spent" in error_msg:
# We held the claim, so nobody else swept it: the client if reclaimed_stale_claim:
# really collected the token. # A prior worker may have redeemed the token and crashed
async with db.create_session() as session: # before finalizing. Treat the ambiguous stale claim as a
await session.exec( # type: ignore[call-overload] # completed sweep rather than misreporting client collection.
update(db.CashuTransaction) await _set_refund_sweep_state(
.where(col(db.CashuTransaction.id) == refund.id) refund.id, swept=True, sweep_started_at=None
.values(collected=True, swept=False) )
else:
await _set_refund_sweep_state(
refund.id,
collected=True,
swept=False,
sweep_started_at=None,
) )
await session.commit()
logger.info( logger.info(
"Refund already spent (client collected), marking swept", "Refund token was already spent",
extra={"id": refund.id},
)
else:
# Release the claim so the next sweep retries this refund.
async with db.create_session() as session:
await session.exec( # type: ignore[call-overload]
update(db.CashuTransaction)
.where(col(db.CashuTransaction.id) == refund.id)
.values(swept=False)
)
await session.commit()
logger.warning(
"Failed to sweep refund",
extra={ extra={
"id": refund.id, "id": refund.id,
"error": str(e), "reclaimed_stale_claim": reclaimed_stale_claim,
}, },
) )
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)
)
if not isinstance(e, Exception):
raise
logger.warning(
"Failed to sweep refund",
extra={"id": refund.id, "error": str(e)},
)
async def refund_sweep_once() -> None: async def refund_sweep_once() -> None:
@@ -1112,15 +1149,21 @@ async def periodic_routstr_fee_payout() -> None:
if accumulated_sats < ROUTSTR_FEE_DEFAULT_PAYOUT: if accumulated_sats < ROUTSTR_FEE_DEFAULT_PAYOUT:
continue continue
paid_msats = accumulated_sats * 1000 paid_msats = accumulated_sats * 1000
payout_checkpointed = await db.reset_routstr_fee(session, paid_msats)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
# Wallet/proof preparation cannot send funds, so do it before the
# durable checkpoint. A preparation failure must not strand an
# in-progress payout that requires manual reconciliation.
wallet = await get_wallet(settings.primary_mint, "sat") wallet = await get_wallet(settings.primary_mint, "sat")
proofs = get_proofs_per_mint_and_unit( proofs = get_proofs_per_mint_and_unit(
wallet, settings.primary_mint, "sat", not_reserved=True wallet, settings.primary_mint, "sat", not_reserved=True
) )
async with db.create_session() as session:
payout_checkpointed = await db.reset_routstr_fee(session, paid_msats)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
try: try:
amount_received = await raw_send_to_lnurl( amount_received = await raw_send_to_lnurl(
wallet, wallet,
@@ -0,0 +1,100 @@
"""Real-DB coverage for db.balances_by_mint_and_unit.
Verifies the grouped liability query used by fetch_all_balances: it sums
balances per (mint_url, unit), filters to the requested mints/units, excludes
NULL mint/currency rows, and returns nothing for empty inputs.
"""
from typing import AsyncGenerator
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, balances_by_mint_and_unit
def _make_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def session() -> "AsyncGenerator[AsyncSession, None]":
engine = _make_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
db_session = AsyncSession(engine, expire_on_commit=False)
try:
yield db_session
finally:
await db_session.close()
await engine.dispose()
async def _add_key(
session: AsyncSession,
hashed_key: str,
balance: int,
mint_url: str | None,
currency: str | None,
) -> None:
session.add(
ApiKey(
hashed_key=hashed_key,
balance=balance,
refund_mint_url=mint_url,
refund_currency=currency,
)
)
await session.commit()
@pytest.mark.asyncio
async def test_sums_and_groups_by_mint_and_unit(session: AsyncSession) -> None:
await _add_key(session, "a", 1000, "http://m1", "sat")
await _add_key(session, "b", 500, "http://m1", "sat")
await _add_key(session, "c", 7000, "http://m1", "msat")
await _add_key(session, "d", 200, "http://m2", "sat")
result = await balances_by_mint_and_unit(
session, ["http://m1", "http://m2"], ["sat", "msat"]
)
assert result[("http://m1", "sat")] == 1500
assert result[("http://m1", "msat")] == 7000
assert result[("http://m2", "sat")] == 200
@pytest.mark.asyncio
async def test_filters_out_unrequested_mints_and_units(session: AsyncSession) -> None:
await _add_key(session, "a", 1000, "http://wanted", "sat")
await _add_key(session, "b", 999, "http://other", "sat")
await _add_key(session, "c", 888, "http://wanted", "usd")
result = await balances_by_mint_and_unit(session, ["http://wanted"], ["sat"])
assert result == {("http://wanted", "sat"): 1000}
@pytest.mark.asyncio
async def test_excludes_rows_with_null_mint_or_currency(session: AsyncSession) -> None:
await _add_key(session, "a", 1000, "http://m1", "sat")
await _add_key(session, "b", 4242, None, None)
result = await balances_by_mint_and_unit(session, ["http://m1"], ["sat"])
assert result == {("http://m1", "sat"): 1000}
@pytest.mark.asyncio
async def test_empty_inputs_return_empty_mapping(session: AsyncSession) -> None:
await _add_key(session, "a", 1000, "http://m1", "sat")
assert await balances_by_mint_and_unit(session, [], ["sat"]) == {}
assert await balances_by_mint_and_unit(session, ["http://m1"], []) == {}
+34 -58
View File
@@ -2,63 +2,45 @@ from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
from sqlalchemy.pool import StaticPool
from routstr.core import db from routstr.core import db
from routstr.core.db import _engine_options from routstr.core.db import create_db_engine
from routstr.core.settings import settings
def test_engine_options_use_bounded_fail_fast_pool(monkeypatch) -> None: # type: ignore[no-untyped-def] @pytest.mark.asyncio
monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False) async def test_engine_uses_validated_bounded_pool_settings(
monkeypatch.delenv("DATABASE_MAX_OVERFLOW", raising=False) monkeypatch: pytest.MonkeyPatch, tmp_path: object
monkeypatch.delenv("DATABASE_POOL_TIMEOUT", raising=False)
monkeypatch.delenv("DATABASE_POOL_RECYCLE", raising=False)
monkeypatch.delenv("DATABASE_POOL_PRE_PING", raising=False)
options = _engine_options("postgresql+asyncpg://db/routstr")
assert options == {
"pool_size": 5,
"max_overflow": 0,
"pool_timeout": 5.0,
"pool_recycle": 1800,
"pool_pre_ping": True,
}
def test_engine_options_are_operator_configurable(monkeypatch) -> None: # type: ignore[no-untyped-def]
monkeypatch.setenv("DATABASE_POOL_SIZE", "12")
monkeypatch.setenv("DATABASE_MAX_OVERFLOW", "3")
monkeypatch.setenv("DATABASE_POOL_TIMEOUT", "2.5")
monkeypatch.setenv("DATABASE_POOL_RECYCLE", "900")
monkeypatch.setenv("DATABASE_POOL_PRE_PING", "false")
assert _engine_options("sqlite+aiosqlite:///keys.db") == {
"pool_size": 12,
"max_overflow": 3,
"pool_timeout": 2.5,
"pool_recycle": 900,
"pool_pre_ping": False,
}
@pytest.mark.parametrize(
("name", "value"),
[
("DATABASE_POOL_SIZE", "0"),
("DATABASE_MAX_OVERFLOW", "-1"),
("DATABASE_POOL_TIMEOUT", "0"),
("DATABASE_POOL_RECYCLE", "-1"),
("DATABASE_POOL_PRE_PING", "maybe"),
],
)
def test_engine_options_reject_invalid_values(
monkeypatch: pytest.MonkeyPatch,
name: str,
value: str,
) -> None: ) -> None:
monkeypatch.setenv(name, value) monkeypatch.setattr(settings, "database_pool_size", 12)
with pytest.raises(ValueError): monkeypatch.setattr(settings, "database_max_overflow", 3)
_engine_options("postgresql+asyncpg://db/routstr") monkeypatch.setattr(settings, "database_pool_timeout", 2.5)
monkeypatch.setattr(settings, "database_pool_recycle", 900)
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
engine = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/pool.db")
try:
assert engine.pool.size() == 12 # type: ignore[attr-defined]
assert engine.pool._max_overflow == 3 # type: ignore[attr-defined]
assert engine.pool._timeout == 2.5 # type: ignore[attr-defined]
assert engine.pool._recycle == 900
assert engine.pool._pre_ping is False
finally:
await engine.dispose()
@pytest.mark.asyncio
async def test_memory_sqlite_keeps_static_pool(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "database_pool_pre_ping", True)
engine = create_db_engine("sqlite+aiosqlite://")
try:
assert isinstance(engine.pool, StaticPool)
assert engine.pool._pre_ping is True
finally:
await engine.dispose()
def test_pool_observer_warns_when_connection_is_held_too_long( def test_pool_observer_warns_when_connection_is_held_too_long(
@@ -75,9 +57,3 @@ def test_pool_observer_warns_when_connection_is_held_too_long(
warning.assert_called_once() warning.assert_called_once()
assert warning.call_args.kwargs["extra"]["held_seconds"] == 12.5 assert warning.call_args.kwargs["extra"]["held_seconds"] == 12.5
def test_memory_sqlite_keeps_dialect_static_pool(monkeypatch) -> None: # type: ignore[no-untyped-def]
monkeypatch.delenv("DATABASE_POOL_PRE_PING", raising=False)
assert _engine_options("sqlite+aiosqlite://") == {"pool_pre_ping": True}
+79 -3
View File
@@ -57,7 +57,7 @@ async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_fee_payout_checkpoints_before_sending() -> None: async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> None:
session = Mock() session = Mock()
fee = SimpleNamespace( fee = SimpleNamespace(
accumulated_msats=5_000, accumulated_msats=5_000,
@@ -67,6 +67,10 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
payout_wallet = Mock() payout_wallet = Mock()
events: list[str] = [] events: list[str] = []
async def prepare(*_args: object) -> Mock:
events.append("prepare")
return payout_wallet
async def checkpoint(*_args: object) -> bool: async def checkpoint(*_args: object) -> bool:
events.append("checkpoint") events.append("checkpoint")
return True return True
@@ -93,14 +97,86 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)), patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint), patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint),
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete), patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)), patch("routstr.wallet.get_wallet", AsyncMock(side_effect=prepare)),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]), patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send), patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
): ):
with pytest.raises(asyncio.CancelledError): with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout() await wallet.periodic_routstr_fee_payout()
assert events == ["checkpoint", "send", "complete"] assert events == ["prepare", "checkpoint", "send", "complete"]
@pytest.mark.asyncio
async def test_fee_payout_preparation_failure_does_not_checkpoint() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
checkpoint = AsyncMock()
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch(
"routstr.wallet.db.create_session", return_value=_session_context(session)
),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", checkpoint),
patch(
"routstr.wallet.get_wallet",
AsyncMock(side_effect=RuntimeError("wallet unavailable")),
),
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
checkpoint.assert_not_awaited()
@pytest.mark.asyncio
async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
send = AsyncMock()
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch(
"routstr.wallet.db.create_session", return_value=_session_context(session)
),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch(
"routstr.wallet.db.reset_routstr_fee",
AsyncMock(return_value=False),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch("routstr.wallet.raw_send_to_lnurl", send),
patch("routstr.wallet.logger.warning") as warning,
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
send.assert_not_awaited()
warning.assert_called_once_with("Routstr fee payout was already claimed")
@pytest.mark.asyncio @pytest.mark.asyncio
+55
View File
@@ -254,3 +254,58 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None:
assert [d["mint_url"] for d in details] == ["http://primary:3338"] assert [d["mint_url"] for d in details] == ["http://primary:3338"]
assert total_wallet == 1000 assert total_wallet == 1000
@pytest.mark.asyncio
async def test_fetch_all_balances_degrades_when_liability_read_fails() -> None:
from routstr.core.settings import settings
with (
patch.object(settings, "cashu_mints", []),
patch.object(settings, "primary_mint", "http://primary:3338"),
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet.db.balances_by_mint_and_unit",
AsyncMock(side_effect=RuntimeError("db pool exhausted")),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
MagicMock(return_value=[MagicMock(amount=1000)]),
),
patch(
"routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs),
),
):
details, total_wallet, total_user, owner = await fetch_all_balances(
units=["sat"]
)
assert details[0]["error"] == "db pool exhausted"
assert details[0]["wallet_balance"] == 1000
assert details[0]["user_balance"] == 0
assert details[0]["owner_balance"] == 0
assert (total_wallet, total_user, owner) == (1000, 0, 0)
@pytest.mark.asyncio
async def test_liability_error_keeps_more_specific_mint_error() -> None:
from routstr.core.settings import settings
with (
patch.object(settings, "cashu_mints", []),
patch.object(settings, "primary_mint", "http://primary:3338"),
patch("routstr.wallet.db.create_session", _fake_session),
patch(
"routstr.wallet.db.balances_by_mint_and_unit",
AsyncMock(side_effect=RuntimeError("db pool exhausted")),
),
patch(
"routstr.wallet.get_wallet",
AsyncMock(side_effect=RuntimeError("mint down")),
),
):
details, *_ = await fetch_all_balances(units=["sat"])
assert details[0]["error"] == "mint down"
+6 -4
View File
@@ -77,7 +77,8 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
), ),
patch( patch(
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0) "routstr.wallet.db.balances_by_mint_and_unit",
AsyncMock(return_value={}),
), ),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
): ):
@@ -130,8 +131,8 @@ async def test_periodic_payout_releases_session_before_slow_mint_send() -> None:
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
), ),
patch( patch(
"routstr.wallet.db.balances_for_mint_and_unit", "routstr.wallet.db.balances_by_mint_and_unit",
AsyncMock(return_value=0), AsyncMock(return_value={}),
), ),
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)), patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)),
): ):
@@ -172,7 +173,8 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
), ),
patch( patch(
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0) "routstr.wallet.db.balances_by_mint_and_unit",
AsyncMock(return_value={}),
), ),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
): ):
+71
View File
@@ -1,3 +1,4 @@
import asyncio
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
@@ -154,3 +155,73 @@ async def test_refund_sweep_records_terminal_but_not_transient_failures(
refund = (await _load(session_factory))["refund"] refund = (await _load(session_factory))["refund"]
assert refund.collected is collected assert refund.collected is collected
assert refund.swept is False assert refund.swept is False
assert refund.sweep_started_at is None
@pytest.mark.asyncio
async def test_refund_sweep_releases_claim_on_cancellation(
session_factory: async_sessionmaker[AsyncSession],
) -> None:
await _insert(
session_factory,
CashuTransaction(
token="cancelled", amount=1, unit="sat", type="out", created_at=800
),
)
with (
patch("routstr.wallet.db.create_session", side_effect=session_factory),
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
patch("routstr.wallet.time.time", return_value=1000),
patch(
"routstr.wallet.recieve_token",
AsyncMock(side_effect=asyncio.CancelledError()),
),
):
with pytest.raises(asyncio.CancelledError):
await refund_sweep_once()
refund = (await _load(session_factory))["cancelled"]
assert refund.swept is False
assert refund.sweep_started_at is None
@pytest.mark.asyncio
async def test_refund_sweep_recovers_stale_claim_without_misreporting_collection(
session_factory: async_sessionmaker[AsyncSession],
) -> None:
await _insert(
session_factory,
CashuTransaction(
token="stale",
amount=1,
unit="sat",
type="out",
created_at=800,
sweep_started_at=100,
),
CashuTransaction(
token="active",
amount=1,
unit="sat",
type="out",
created_at=800,
sweep_started_at=950,
),
)
receive = AsyncMock(side_effect=RuntimeError("token already spent"))
with (
patch("routstr.wallet.db.create_session", side_effect=session_factory),
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
patch("routstr.wallet.time.time", return_value=1000),
patch("routstr.wallet.recieve_token", receive),
):
await refund_sweep_once()
receive.assert_awaited_once_with("stale")
loaded = await _load(session_factory)
assert loaded["stale"].swept is True
assert loaded["stale"].collected is False
assert loaded["stale"].sweep_started_at is None
assert loaded["active"].swept is False
assert loaded["active"].sweep_started_at == 950
+95
View File
@@ -62,6 +62,101 @@ def test_payout_settings_have_sensible_defaults() -> None:
assert s.payout_interval_seconds == 900 assert s.payout_interval_seconds == 900
def test_database_pool_defaults_are_bounded_and_fail_fast() -> 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_pool_recycle == 1800
assert s.database_pool_pre_ping is True
assert s.database_pool_hold_warn_seconds == 10.0
@pytest.mark.parametrize(
("field", "bad_value"),
[
("database_pool_size", 0),
("database_max_overflow", -1),
("database_pool_timeout", 0),
("database_pool_recycle", -1),
("database_pool_hold_warn_seconds", 0),
],
)
def test_database_pool_settings_reject_invalid_values(
field: str, bad_value: int
) -> None:
with pytest.raises(ValidationError):
Settings(**{field: bad_value})
@pytest.mark.asyncio
async def test_database_pool_fields_are_env_only_not_persisted(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""DB pool sizing is infrastructure the node needs *before* it can read the
DB, so it can never be configured from the DB — it must never be written to
the settings blob, and a stale/injected DB value must never shadow env.
"""
monkeypatch.setenv("DATABASE_POOL_SIZE", "7")
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with AsyncSession(engine, expire_on_commit=False) as session:
s = await SettingsService.initialize(session)
# The env value is live for runtime consumers...
assert s.database_pool_size == 7
# ...but pool sizing is never written to the settings blob.
blob = await _read_settings_blob(session)
for field in (
"database_pool_size",
"database_max_overflow",
"database_pool_timeout",
"database_pool_recycle",
"database_pool_pre_ping",
"database_pool_hold_warn_seconds",
):
assert field not in blob
# Even a stale blob that somehow carries a pool value must not win: env
# stays authoritative on the next initialize.
await session.exec( # type: ignore
text("UPDATE settings SET data = :d WHERE id = 1").bindparams(
d=json.dumps({"database_pool_size": 99})
)
)
await session.commit()
again = await SettingsService.initialize(session)
assert again.database_pool_size == 7
@pytest.mark.asyncio
async def test_update_does_not_apply_env_only_fields_to_live_settings(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""DB pool sizing is env-only: a settings update must neither persist it nor
mutate the live value. The engine pool is already built at boot from env, so
a UI/API update carrying a pool value must not make the live setting diverge
from the running pool.
"""
monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False)
monkeypatch.setattr(settings, "database_pool_size", 5)
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with AsyncSession(engine, expire_on_commit=False) as session:
await SettingsService.initialize(session)
await SettingsService.update(
{"database_pool_size": 99, "name": "PoolTweaker"}, session
)
# A non-env-only field still updates normally...
assert settings.name == "PoolTweaker"
# ...but the env-only pool size stays at the boot value.
assert settings.database_pool_size == 5
# ...and it is never written to the settings blob.
blob = await _read_settings_blob(session)
assert "database_pool_size" not in blob
@pytest.mark.parametrize( @pytest.mark.parametrize(
"field,bad_value", "field,bad_value",
[ [