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
@@ -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
import pytest
from sqlalchemy.pool import StaticPool
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]
monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False)
monkeypatch.delenv("DATABASE_MAX_OVERFLOW", raising=False)
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,
@pytest.mark.asyncio
async def test_engine_uses_validated_bounded_pool_settings(
monkeypatch: pytest.MonkeyPatch, tmp_path: object
) -> None:
monkeypatch.setenv(name, value)
with pytest.raises(ValueError):
_engine_options("postgresql+asyncpg://db/routstr")
monkeypatch.setattr(settings, "database_pool_size", 12)
monkeypatch.setattr(settings, "database_max_overflow", 3)
monkeypatch.setattr(settings, "database_pool_timeout", 2.5)
monkeypatch.setattr(settings, "database_pool_recycle", 900)
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
engine = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/pool.db")
try:
assert engine.pool.size() == 12 # type: ignore[attr-defined]
assert engine.pool._max_overflow == 3 # type: ignore[attr-defined]
assert engine.pool._timeout == 2.5 # type: ignore[attr-defined]
assert engine.pool._recycle == 900
assert engine.pool._pre_ping is False
finally:
await engine.dispose()
@pytest.mark.asyncio
async def test_memory_sqlite_keeps_static_pool(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "database_pool_pre_ping", True)
engine = create_db_engine("sqlite+aiosqlite://")
try:
assert isinstance(engine.pool, StaticPool)
assert engine.pool._pre_ping is True
finally:
await engine.dispose()
def test_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()
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
async def test_fee_payout_checkpoints_before_sending() -> None:
async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
@@ -67,6 +67,10 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
payout_wallet = Mock()
events: list[str] = []
async def prepare(*_args: object) -> Mock:
events.append("prepare")
return payout_wallet
async def checkpoint(*_args: object) -> bool:
events.append("checkpoint")
return True
@@ -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.reset_routstr_fee", side_effect=checkpoint),
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)),
patch("routstr.wallet.get_wallet", AsyncMock(side_effect=prepare)),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
assert events == ["checkpoint", "send", "complete"]
assert events == ["prepare", "checkpoint", "send", "complete"]
@pytest.mark.asyncio
async def test_fee_payout_preparation_failure_does_not_checkpoint() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
checkpoint = AsyncMock()
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch(
"routstr.wallet.db.create_session", return_value=_session_context(session)
),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch("routstr.wallet.db.reset_routstr_fee", checkpoint),
patch(
"routstr.wallet.get_wallet",
AsyncMock(side_effect=RuntimeError("wallet unavailable")),
),
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
checkpoint.assert_not_awaited()
@pytest.mark.asyncio
async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
session = Mock()
fee = SimpleNamespace(
accumulated_msats=5_000,
payout_in_progress_msats=0,
payout_started_at=None,
)
send = AsyncMock()
with (
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
patch(
"routstr.wallet.asyncio.sleep",
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
),
patch(
"routstr.wallet.db.create_session", return_value=_session_context(session)
),
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
patch(
"routstr.wallet.db.reset_routstr_fee",
AsyncMock(return_value=False),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
patch("routstr.wallet.raw_send_to_lnurl", send),
patch("routstr.wallet.logger.warning") as warning,
):
with pytest.raises(asyncio.CancelledError):
await wallet.periodic_routstr_fee_payout()
send.assert_not_awaited()
warning.assert_called_once_with("Routstr fee payout was already claimed")
@pytest.mark.asyncio
+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 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),
),
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),
):
@@ -130,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_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", 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),
),
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),
):
+71
View File
@@ -1,3 +1,4 @@
import asyncio
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
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"]
assert refund.collected is collected
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
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(
"field,bad_value",
[