From bb40e869e63348917a1e871ab04ea16ad2572b86 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 17:32:31 +0200 Subject: [PATCH 1/4] fix: bound owner payout by per-mint liability instead of total --- routstr/core/db.py | 29 +++++ routstr/wallet.py | 78 +++++++++--- tests/unit/test_lnurl_change.py | 4 + tests/unit/test_payout_liability_bounds.py | 140 +++++++++++++++++++++ tests/unit/test_payout_limits.py | 13 +- tests/unit/test_periodic_payout.py | 32 ++++- 6 files changed, 272 insertions(+), 24 deletions(-) create mode 100644 tests/unit/test_payout_liability_bounds.py diff --git a/routstr/core/db.py b/routstr/core/db.py index 78019c32..057513dd 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -1098,6 +1098,35 @@ async def total_user_liability(db_session: AsyncSession) -> int: return int(result.one() or 0) +async def user_liability_for_mint_and_unit( + db_session: AsyncSession, mint_url: str, unit: str +) -> int: + """Return outstanding user funds that refund from one mint and unit, in msats. + + Key balances and unresolved refund claims are summed in one statement for + the same reason as ``total_user_liability``. + """ + key_balances = ( + select(func.coalesce(func.sum(ApiKey.balance), 0)) + .where( + col(ApiKey.refund_mint_url) == mint_url, + col(ApiKey.refund_currency) == unit, + ) + .scalar_subquery() + ) + unresolved_refunds = ( + select(func.coalesce(func.sum(Refund.amount_msats), 0)) + .where( + col(Refund.status).in_(REFUND_UNRESOLVED_STATUSES), + col(Refund.mint_url) == mint_url, + col(Refund.unit) == unit, + ) + .scalar_subquery() + ) + result = await db_session.exec(select(key_balances + unresolved_refunds)) + return int(result.one() or 0) + + async def balance_for_mint_and_unit( db_session: AsyncSession, mint_url: str, unit: str ) -> int: diff --git a/routstr/wallet.py b/routstr/wallet.py index 2c592b3b..7054d38d 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -701,18 +701,65 @@ class Bolt11PaymentPlan: return maximum if self.unit == "sat" else (maximum + 999) // 1000 +def _to_msats(amount: int, unit: str) -> int: + return _sats_to_msats(amount) if unit == "sat" else amount + + +async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int: + """Sum unreserved proofs of every other trusted wallet, in msats. + + Reads local proof snapshots only. A wallet that cannot be loaded counts + as empty, which can only shrink the owner surplus. + """ + total = 0 + for other_mint in _mints_to_inspect(): + for other_unit in ("sat", "msat"): + if (other_mint, other_unit) == (mint_url, unit): + continue + try: + wallet = await get_wallet(other_mint, other_unit) + except Exception as e: + logger.debug( + "Wallet excluded from owner surplus", + extra={ + "mint_url": other_mint, + "unit": other_unit, + "error": str(e), + }, + ) + continue + proofs = get_proofs_per_mint_and_unit( + wallet, other_mint, other_unit, not_reserved=True + ) + total += _to_msats(sum(proof.amount for proof in proofs), other_unit) + return total + + async def _owner_balance_for_mint_and_unit( mint_url: str, unit: str, proofs_balance: int ) -> int: - """Return spendable node-owned funds without crossing user liabilities.""" + """Return owner funds in one wallet, in that wallet's unit, never negative. + + A key's refund mint is a preference, not funding provenance: a key topped + up from a second mint keeps its original refund mint. So two bounds apply. + The wallet keeps the liability declared against it, so refunds drawn from + it stay serviceable. All wallets together keep the total liability, so + misattributed customer funds are never paid out as profit. + """ + others_msats = await _other_wallets_unreserved_msats(mint_url, unit) async with db.create_session() as session: - # Refund mint is a preference, not funding provenance. Mirror payout's - # conservative rule and protect the full liability at every mint. - user_liability = await db.total_user_liability(session) - # API-key balances are stored in msats. Cashu ``sat`` proofs are not. - if unit == "sat": - user_liability = _msats_to_sats_ceil(user_liability) - return max(0, proofs_balance - user_liability) + mint_liability = await db.user_liability_for_mint_and_unit( + session, mint_url, unit + ) + total_liability = await db.total_user_liability(session) + proofs_msats = _to_msats(proofs_balance, unit) + surplus_msats = min( + proofs_msats - mint_liability, + proofs_msats + others_msats - total_liability, + ) + # Cashu ``sat`` proofs are whole sats; round the surplus down, never up. + surplus = _msats_to_sats(surplus_msats) if unit == "sat" else surplus_msats + return max(0, surplus) async def maximum_owner_cashu_balance_sats() -> int: @@ -1650,15 +1697,12 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: ) return - # Fetch liability after the proofs snapshot and settle delay while the + # Read liabilities after the proofs snapshot and settle delay while the # wallet operation guard excludes concurrent proof mutation and crediting. try: - async with db.create_session() as session: - # ApiKey stores a refund preference, not funding provenance. Until - # liabilities have a durable per-credit ledger, subtract the total - # liability from every wallet rather than risk calling customer - # funds owner profit on the wrong mint. - user_balance = await db.total_user_liability(session) + available_balance = await _owner_balance_for_mint_and_unit( + mint_url, unit, sum(proof.amount for proof in proofs) + ) except Exception as e: logger.error( f"Error in periodic payout cycle: {type(e).__name__}", @@ -1667,10 +1711,6 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: return try: - if unit == "sat": - user_balance = _msats_to_sats_ceil(user_balance) - proofs_balance = sum(proof.amount for proof in proofs) - available_balance = proofs_balance - user_balance max_amount = ( settings.max_payout_sat if unit == "sat" diff --git a/tests/unit/test_lnurl_change.py b/tests/unit/test_lnurl_change.py index b7566247..fc62d02c 100644 --- a/tests/unit/test_lnurl_change.py +++ b/tests/unit/test_lnurl_change.py @@ -129,6 +129,10 @@ async def test_capped_payout_recovers_all_change_with_real_cashu_sdk( "routstr.wallet.db.total_user_liability", AsyncMock(return_value=liability * 1000), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=liability * 1000), + ), patch( "routstr.payment.lnurl.get_lnurl_data", AsyncMock( diff --git a/tests/unit/test_payout_liability_bounds.py b/tests/unit/test_payout_liability_bounds.py new file mode 100644 index 00000000..f54c5a0c --- /dev/null +++ b/tests/unit/test_payout_liability_bounds.py @@ -0,0 +1,140 @@ +"""Owner payout keeps each wallet's declared liability and the global total. + +Regression for multi-mint payout starvation: subtracting the *total* user +liability from every wallet hid the owner surplus on any mint holding less +than the whole liability, so only the largest wallet could ever pay out. +""" + +from collections.abc import AsyncIterator, Iterator +from contextlib import ExitStack, asynccontextmanager, contextmanager +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +from routstr.core.settings import settings +from routstr.wallet import _owner_balance_for_mint_and_unit, _payout_mint_and_unit + +MINT_A = "https://a.test" +MINT_B = "https://b.test" + + +@asynccontextmanager +async def _session() -> AsyncIterator[Mock]: + yield Mock() + + +def _wallets( + sat_proofs: dict[str, int], unreachable: frozenset[str] = frozenset() +) -> tuple[AsyncMock, Mock]: + """Fake get_wallet/get_proofs for sat wallets; msat wallets are unsupported.""" + + async def get_wallet(mint_url: str, unit: str, **_: object) -> Mock: + if unit != "sat" or mint_url in unreachable: + raise ValueError("unsupported") + return Mock(url=mint_url) + + def get_proofs(wallet: Mock, mint_url: str, unit: str, **_: object) -> list[Mock]: + return [Mock(amount=sat_proofs[mint_url])] + + return AsyncMock(side_effect=get_wallet), Mock(side_effect=get_proofs) + + +@contextmanager +def _liabilities(per_mint_sats: dict[str, int], total_sats: int) -> Iterator[None]: + async def per_mint(_session: object, mint_url: str, unit: str) -> int: + return per_mint_sats.get(mint_url, 0) * 1000 + + with ExitStack() as stack: + for target in ( + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(side_effect=per_mint), + ), + patch( + "routstr.wallet.db.total_user_liability", + AsyncMock(return_value=total_sats * 1000), + ), + patch("routstr.wallet.db.create_session", _session), + patch.object(settings, "cashu_mints", [MINT_A, MINT_B]), + patch.object(settings, "primary_mint", MINT_A), + ): + stack.enter_context(target) + yield + + +@pytest.mark.asyncio +async def test_owner_balance_keeps_only_the_wallets_own_liability() -> None: + """Mint B's surplus is bounded by B's liability, not by A's.""" + get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) + with ( + _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + ): + assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 236 + assert await _owner_balance_for_mint_and_unit(MINT_A, "sat", 400) == 184 + + +@pytest.mark.asyncio +async def test_owner_balance_never_exceeds_global_surplus() -> None: + """Liability nobody declared against a mint is still covered in aggregate.""" + get_wallet, get_proofs = _wallets({MINT_A: 100, MINT_B: 270}) + with ( + _liabilities({}, total_sats=250), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + ): + assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 120 + + +@pytest.mark.asyncio +async def test_unloadable_wallet_counts_as_empty() -> None: + """A wallet that cannot be read shrinks the surplus rather than inflating it.""" + get_wallet, get_proofs = _wallets( + {MINT_A: 400, MINT_B: 270}, unreachable=frozenset({MINT_A}) + ) + with ( + _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + ): + assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 20 + + +@pytest.mark.asyncio +async def test_msat_wallet_surplus_is_not_rounded() -> None: + get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0}) + with ( + _liabilities({MINT_B: 0}, total_sats=0), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=1_500), + ), + patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=1_500)), + ): + assert await _owner_balance_for_mint_and_unit(MINT_B, "msat", 4_000) == 2_500 + + +@pytest.mark.asyncio +async def test_payout_sends_the_smaller_wallets_surplus() -> None: + """End to end: the wallet below the total liability still pays out.""" + get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) + send = AsyncMock(return_value=236_000) + with ( + _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), + patch.object(settings, "min_payout_sat", 50), + patch.object(settings, "max_payout_sat", 250_000), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + patch( + "routstr.wallet.slow_filter_spend_proofs", + AsyncMock(side_effect=lambda proofs, wallet: proofs), + ), + patch("routstr.wallet.asyncio.sleep", AsyncMock()), + patch("routstr.wallet.raw_send_to_lnurl", send), + ): + await _payout_mint_and_unit(MINT_B, "sat") + assert send.await_args is not None + assert send.await_args.kwargs["amount"] == 236 diff --git a/tests/unit/test_payout_limits.py b/tests/unit/test_payout_limits.py index 8d3df75e..48137a25 100644 --- a/tests/unit/test_payout_limits.py +++ b/tests/unit/test_payout_limits.py @@ -1,6 +1,6 @@ from collections.abc import AsyncIterator from contextlib import asynccontextmanager -from unittest.mock import AsyncMock, Mock, patch +from unittest.mock import AsyncMock, Mock, call, patch import pytest @@ -39,13 +39,18 @@ async def test_payout_limits_and_proof_refresh( patch( "routstr.wallet.db.total_user_liability", AsyncMock(return_value=liability) ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=liability), + ), patch("routstr.wallet.asyncio.sleep", sleep), patch("routstr.wallet.raw_send_to_lnurl", send), ): await _payout_mint_and_unit("https://mint.test", unit) - get_wallet.assert_awaited_once_with( - "https://mint.test", unit, force_reload_proofs=True - ) + reloads = [ + c for c in get_wallet.await_args_list if c.kwargs.get("force_reload_proofs") + ] + assert reloads == [call("https://mint.test", unit, force_reload_proofs=True)] if expected is None: send.assert_not_awaited() else: diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index 03871a82..ce5f0014 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -107,6 +107,10 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non "routstr.wallet.db.total_user_liability", AsyncMock(return_value=0), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch("routstr.wallet.db.record_lightning_payout", record_payout), patch("routstr.wallet.db.settle_lightning_payout", settle_payout), patch("routstr.wallet.raw_send_to_lnurl", raw_send), @@ -181,6 +185,10 @@ async def test_periodic_payout_releases_session_before_slow_mint_send() -> None: "routstr.wallet.db.total_user_liability", AsyncMock(return_value=0), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)), ): with pytest.raises(_LoopBreak): @@ -229,6 +237,10 @@ async def test_periodic_payout_isolates_failing_mint() -> None: "routstr.wallet.db.total_user_liability", AsyncMock(return_value=0), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch("routstr.wallet.raw_send_to_lnurl", raw_send), ): with pytest.raises(_LoopBreak): @@ -237,7 +249,9 @@ async def test_periodic_payout_isolates_failing_mint() -> None: # The bad mint raised on get_wallet for both units, yet the good mint was # still reached and paid out for both units — failures are isolated. good_calls = [ - c for c in get_wallet.await_args_list if c.args[0] == "http://good:3338" + c + for c in get_wallet.await_args_list + if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs") ] assert len(good_calls) == 2 # sat + msat assert raw_send.await_count == 2 # good mint paid for both units @@ -331,6 +345,10 @@ async def test_periodic_payout_caps_amount_at_max_payout_sat() -> None: "routstr.wallet.db.total_user_liability", AsyncMock(return_value=0), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch("routstr.wallet.raw_send_to_lnurl", raw_send), ): with pytest.raises(_LoopBreak): @@ -379,6 +397,10 @@ async def test_payout_history_records_the_capped_amount() -> None: AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch( "routstr.wallet.db.list_unsettled_lightning_payouts", AsyncMock(return_value=[]), @@ -446,6 +468,10 @@ async def test_payout_history_marks_failed_only_on_proven_non_payment( AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch( "routstr.wallet.db.list_unsettled_lightning_payouts", AsyncMock(return_value=[]), @@ -503,6 +529,10 @@ async def test_payout_history_write_failure_does_not_block_payout() -> None: AsyncMock(side_effect=lambda proofs, wallet: proofs), ), patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=0), + ), patch( "routstr.wallet.db.list_unsettled_lightning_payouts", AsyncMock(return_value=[]), From b663b1d83ae992092f30e2f3a01011cd8c62692c Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 19:42:49 +0200 Subject: [PATCH 2/4] fix: read other wallets' proofs fresh when bounding owner payout --- routstr/core/db.py | 3 +- routstr/wallet.py | 32 ++--- tests/unit/test_payout_liability_bounds.py | 48 ++++++- tests/unit/test_payout_limits.py | 8 +- tests/unit/test_periodic_payout.py | 13 +- .../test_user_liability_for_mint_and_unit.py | 136 ++++++++++++++++++ tests/unit/test_wallet.py | 12 +- 7 files changed, 214 insertions(+), 38 deletions(-) create mode 100644 tests/unit/test_user_liability_for_mint_and_unit.py diff --git a/routstr/core/db.py b/routstr/core/db.py index 057513dd..500420d5 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -1103,8 +1103,7 @@ async def user_liability_for_mint_and_unit( ) -> int: """Return outstanding user funds that refund from one mint and unit, in msats. - Key balances and unresolved refund claims are summed in one statement for - the same reason as ``total_user_liability``. + Single statement, for the same atomicity reason as ``total_user_liability``. """ key_balances = ( select(func.coalesce(func.sum(ApiKey.balance), 0)) diff --git a/routstr/wallet.py b/routstr/wallet.py index 7054d38d..8ade952a 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -121,7 +121,7 @@ def _msats_to_sats_ceil(amount: int) -> int: def _mints_to_inspect() -> list[str]: """Return configured mints plus the primary mint, without duplicates.""" - mint_urls = list(settings.cashu_mints) + mint_urls = list(dict.fromkeys(settings.cashu_mints)) if settings.primary_mint and settings.primary_mint not in mint_urls: mint_urls.append(settings.primary_mint) return mint_urls @@ -708,8 +708,10 @@ def _to_msats(amount: int, unit: str) -> int: async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int: """Sum unreserved proofs of every other trusted wallet, in msats. - Reads local proof snapshots only. A wallet that cannot be loaded counts - as empty, which can only shrink the owner surplus. + This total only ever raises the payout ceiling, so proofs are reloaded from + the local db: a cached snapshot up to 30s stale could still hide another + process's reservation. A wallet that cannot be loaded counts as empty, + which can only shrink the owner surplus. """ total = 0 for other_mint in _mints_to_inspect(): @@ -717,16 +719,11 @@ async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int: if (other_mint, other_unit) == (mint_url, unit): continue try: - wallet = await get_wallet(other_mint, other_unit) - except Exception as e: - logger.debug( - "Wallet excluded from owner surplus", - extra={ - "mint_url": other_mint, - "unit": other_unit, - "error": str(e), - }, + wallet = await get_wallet( + other_mint, other_unit, force_reload_proofs=True ) + except Exception as e: + logger.debug(f"Wallet {other_mint} {other_unit} excluded: {e}") continue proofs = get_proofs_per_mint_and_unit( wallet, other_mint, other_unit, not_reserved=True @@ -738,13 +735,12 @@ async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int: async def _owner_balance_for_mint_and_unit( mint_url: str, unit: str, proofs_balance: int ) -> int: - """Return owner funds in one wallet, in that wallet's unit, never negative. + """Return owner funds in one wallet, in that wallet's unit. A key's refund mint is a preference, not funding provenance: a key topped - up from a second mint keeps its original refund mint. So two bounds apply. - The wallet keeps the liability declared against it, so refunds drawn from - it stay serviceable. All wallets together keep the total liability, so - misattributed customer funds are never paid out as profit. + up from a second mint keeps its original refund mint. Hence two bounds — + the per-mint one keeps refunds serviceable from the mint they name, the + global one stops misattributed customer funds being paid out as profit. """ others_msats = await _other_wallets_unreserved_msats(mint_url, unit) async with db.create_session() as session: @@ -757,7 +753,7 @@ async def _owner_balance_for_mint_and_unit( proofs_msats - mint_liability, proofs_msats + others_msats - total_liability, ) - # Cashu ``sat`` proofs are whole sats; round the surplus down, never up. + # Cashu ``sat`` proofs are whole sats. surplus = _msats_to_sats(surplus_msats) if unit == "sat" else surplus_msats return max(0, surplus) diff --git a/tests/unit/test_payout_liability_bounds.py b/tests/unit/test_payout_liability_bounds.py index f54c5a0c..3bb91e6c 100644 --- a/tests/unit/test_payout_liability_bounds.py +++ b/tests/unit/test_payout_liability_bounds.py @@ -102,24 +102,60 @@ async def test_unloadable_wallet_counts_as_empty() -> None: @pytest.mark.asyncio -async def test_msat_wallet_surplus_is_not_rounded() -> None: +async def test_duplicate_configured_mint_is_counted_once() -> None: + """A mint listed twice in CASHU_MINTS would otherwise raise the global bound.""" + get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) + with ( + _liabilities({}, total_sats=600), + patch.object(settings, "cashu_mints", [MINT_A, MINT_A, MINT_B]), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + ): + assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 70 + + +@pytest.mark.asyncio +async def test_other_wallets_are_read_from_fresh_local_proofs() -> None: + """A stale snapshot of another wallet would raise the global bound.""" + get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) + with ( + _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + ): + await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) + assert get_wallet.await_args_list + assert all(c.kwargs.get("force_reload_proofs") for c in get_wallet.await_args_list) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "mint_liability,total_liability,expected", + [(1_500, 2_200, 1_800), (2_700, 1_500, 1_300)], +) +async def test_msat_wallet_surplus_is_not_rounded( + mint_liability: int, total_liability: int, expected: int +) -> None: + """Either bound can bind, and neither is rounded to whole sats.""" get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0}) with ( - _liabilities({MINT_B: 0}, total_sats=0), + _liabilities({}, total_sats=0), patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), patch( "routstr.wallet.db.user_liability_for_mint_and_unit", - AsyncMock(return_value=1_500), + AsyncMock(return_value=mint_liability), + ), + patch( + "routstr.wallet.db.total_user_liability", + AsyncMock(return_value=total_liability), ), - patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=1_500)), ): - assert await _owner_balance_for_mint_and_unit(MINT_B, "msat", 4_000) == 2_500 + assert await _owner_balance_for_mint_and_unit(MINT_B, "msat", 4_000) == expected @pytest.mark.asyncio async def test_payout_sends_the_smaller_wallets_surplus() -> None: - """End to end: the wallet below the total liability still pays out.""" get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) send = AsyncMock(return_value=236_000) with ( diff --git a/tests/unit/test_payout_limits.py b/tests/unit/test_payout_limits.py index 48137a25..87403074 100644 --- a/tests/unit/test_payout_limits.py +++ b/tests/unit/test_payout_limits.py @@ -47,10 +47,10 @@ async def test_payout_limits_and_proof_refresh( patch("routstr.wallet.raw_send_to_lnurl", send), ): await _payout_mint_and_unit("https://mint.test", unit) - reloads = [ - c for c in get_wallet.await_args_list if c.kwargs.get("force_reload_proofs") - ] - assert reloads == [call("https://mint.test", unit, force_reload_proofs=True)] + # Later awaits belong to the other-wallet scan, which forces a reload too. + assert get_wallet.await_args_list[0] == call( + "https://mint.test", unit, force_reload_proofs=True + ) if expected is None: send.assert_not_awaited() else: diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index ce5f0014..9e57f901 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -14,7 +14,7 @@ from collections.abc import Callable, Coroutine from contextlib import asynccontextmanager from pathlib import Path from typing import Any -from unittest.mock import ANY, AsyncMock, MagicMock, patch +from unittest.mock import ANY, AsyncMock, MagicMock, call, patch import pytest @@ -248,12 +248,11 @@ async def test_periodic_payout_isolates_failing_mint() -> None: # The bad mint raised on get_wallet for both units, yet the good mint was # still reached and paid out for both units — failures are isolated. - good_calls = [ - c - for c in get_wallet.await_args_list - if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs") - ] - assert len(good_calls) == 2 # sat + msat + for unit in ("sat", "msat"): + assert ( + call("http://good:3338", unit, force_reload_proofs=True) + in get_wallet.await_args_list + ) assert raw_send.await_count == 2 # good mint paid for both units diff --git a/tests/unit/test_user_liability_for_mint_and_unit.py b/tests/unit/test_user_liability_for_mint_and_unit.py new file mode 100644 index 00000000..4ce00941 --- /dev/null +++ b/tests/unit/test_user_liability_for_mint_and_unit.py @@ -0,0 +1,136 @@ +"""Real-DB coverage for db.user_liability_for_mint_and_unit. + +Verifies the per-mint liability query that bounds owner payout: it sums key +balances and unresolved refund claims for one (mint_url, unit), excludes +resolved claims and other mints/units, and drops keys with no refund mint. +""" + +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, Refund, user_liability_for_mint_and_unit + +MINT = "http://m1" + + +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 = MINT, + currency: str | None = "sat", +) -> None: + session.add( + ApiKey( + hashed_key=hashed_key, + balance=balance, + refund_mint_url=mint_url, + refund_currency=currency, + ) + ) + await session.commit() + + +async def _add_refund( + session: AsyncSession, + hashed_key: str, + amount_msats: int, + status: str, + mint_url: str = MINT, + unit: str = "sat", +) -> None: + session.add( + Refund( + api_key_hashed_key=hashed_key, + method="lightning", + amount_msats=amount_msats, + unit=unit, + mint_url=mint_url, + status=status, + ) + ) + await session.commit() + + +@pytest.mark.asyncio +async def test_sums_key_balances_for_the_mint_and_unit(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_key(session, "b", 500) + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1500 + + +@pytest.mark.asyncio +async def test_adds_unresolved_refunds_to_key_balances(session: AsyncSession) -> None: + # One open claim per key, so each unresolved status needs its own key. + await _add_key(session, "a", 1000) + await _add_refund(session, "a", 300, "pending") + await _add_key(session, "b", 0) + await _add_refund(session, "b", 40, "ambiguous") + await _add_key(session, "c", 0) + await _add_refund(session, "c", 7, "stuck") + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1347 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("status", ["paid", "failed"]) +async def test_excludes_resolved_refunds(session: AsyncSession, status: str) -> None: + await _add_key(session, "a", 0) + await _add_refund(session, "a", 900, status) + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 0 + + +@pytest.mark.asyncio +async def test_excludes_other_mints_and_units(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_key(session, "other-mint", 111, mint_url="http://m2") + await _add_key(session, "other-unit", 222, currency="msat") + await _add_refund(session, "a", 300, "pending") + await _add_refund(session, "other-mint", 444, "pending", mint_url="http://m2") + await _add_refund(session, "other-unit", 555, "pending", unit="msat") + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1300 + + +@pytest.mark.asyncio +async def test_excludes_keys_without_a_refund_mint(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_key(session, "unattributed", 4242, mint_url=None, currency=None) + + assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1000 + + +@pytest.mark.asyncio +async def test_unknown_mint_has_no_liability(session: AsyncSession) -> None: + await _add_key(session, "a", 1000) + await _add_refund(session, "a", 300, "pending") + + assert await user_liability_for_mint_and_unit(session, "http://missing", "sat") == 0 diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 0d4f7e9e..8054f91c 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -1254,9 +1254,14 @@ async def test_prepare_bolt11_payment_does_not_spend_user_liabilities() -> None: @pytest.mark.asyncio -async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() -> None: +async def test_prepare_bolt11_payment_floors_fractional_owner_surplus() -> None: + """A sub-sat surplus is not enough to fund a 1 sat invoice.""" from routstr.core.settings import settings + @asynccontextmanager + async def session() -> AsyncIterator[MagicMock]: + yield MagicMock() + wallet = MagicMock() wallet.proofs = [MagicMock(amount=100)] wallet.melt_quote = AsyncMock( @@ -1281,10 +1286,15 @@ async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() - "routstr.wallet.slow_filter_spend_proofs", side_effect=lambda proofs, wallet: proofs, ), + patch("routstr.wallet.db.create_session", session), patch( "routstr.wallet.db.total_user_liability", AsyncMock(return_value=99_999), ), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(return_value=99_999), + ), pytest.raises(ValueError, match="user liabilities"), ): await prepare_bolt11_payment("lnbc-invoice") From 90be8e6792790afd9b66d47a676907287b4ebd3b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 19:55:26 +0200 Subject: [PATCH 3/4] refactor: simplify payout liability tests and tighten reload assertion --- routstr/wallet.py | 9 ++-- tests/unit/test_payout_liability_bounds.py | 43 +++++++++++-------- tests/unit/test_periodic_payout.py | 14 +++--- .../test_user_liability_for_mint_and_unit.py | 9 +--- 4 files changed, 38 insertions(+), 37 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index 8ade952a..c69f1aaa 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -830,9 +830,7 @@ async def _prepare_bolt11_payment(invoice: str) -> Bolt11PaymentPlan: ) if owner_balance < required: continue - owner_balance_msats = ( - owner_balance * 1000 if unit == "sat" else owner_balance - ) + owner_balance_msats = _to_msats(owner_balance, unit) candidates.append( (owner_balance_msats, wallet, proofs, quote, mint_url, unit) ) @@ -1693,8 +1691,9 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None: ) return - # Read liabilities after the proofs snapshot and settle delay while the - # wallet operation guard excludes concurrent proof mutation and crediting. + # Read liabilities and the other wallets' proofs after this wallet's proofs + # snapshot and settle delay, while the wallet operation guard excludes + # concurrent proof mutation and crediting. try: available_balance = await _owner_balance_for_mint_and_unit( mint_url, unit, sum(proof.amount for proof in proofs) diff --git a/tests/unit/test_payout_liability_bounds.py b/tests/unit/test_payout_liability_bounds.py index 3bb91e6c..98979e99 100644 --- a/tests/unit/test_payout_liability_bounds.py +++ b/tests/unit/test_payout_liability_bounds.py @@ -6,7 +6,7 @@ than the whole liability, so only the largest wallet could ever pay out. """ from collections.abc import AsyncIterator, Iterator -from contextlib import ExitStack, asynccontextmanager, contextmanager +from contextlib import asynccontextmanager, contextmanager from unittest.mock import AsyncMock, Mock, patch import pytest @@ -39,32 +39,37 @@ def _wallets( return AsyncMock(side_effect=get_wallet), Mock(side_effect=get_proofs) +@contextmanager +def _env() -> Iterator[None]: + with ( + patch("routstr.wallet.db.create_session", _session), + patch.object(settings, "cashu_mints", [MINT_A, MINT_B]), + patch.object(settings, "primary_mint", MINT_A), + ): + yield + + @contextmanager def _liabilities(per_mint_sats: dict[str, int], total_sats: int) -> Iterator[None]: async def per_mint(_session: object, mint_url: str, unit: str) -> int: return per_mint_sats.get(mint_url, 0) * 1000 - with ExitStack() as stack: - for target in ( - patch( - "routstr.wallet.db.user_liability_for_mint_and_unit", - AsyncMock(side_effect=per_mint), - ), - patch( - "routstr.wallet.db.total_user_liability", - AsyncMock(return_value=total_sats * 1000), - ), - patch("routstr.wallet.db.create_session", _session), - patch.object(settings, "cashu_mints", [MINT_A, MINT_B]), - patch.object(settings, "primary_mint", MINT_A), - ): - stack.enter_context(target) + with ( + _env(), + patch( + "routstr.wallet.db.user_liability_for_mint_and_unit", + AsyncMock(side_effect=per_mint), + ), + patch( + "routstr.wallet.db.total_user_liability", + AsyncMock(return_value=total_sats * 1000), + ), + ): yield @pytest.mark.asyncio async def test_owner_balance_keeps_only_the_wallets_own_liability() -> None: - """Mint B's surplus is bounded by B's liability, not by A's.""" get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) with ( _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), @@ -89,7 +94,7 @@ async def test_owner_balance_never_exceeds_global_surplus() -> None: @pytest.mark.asyncio async def test_unloadable_wallet_counts_as_empty() -> None: - """A wallet that cannot be read shrinks the surplus rather than inflating it.""" + """Shrinks the surplus rather than inflating it.""" get_wallet, get_proofs = _wallets( {MINT_A: 400, MINT_B: 270}, unreachable=frozenset({MINT_A}) ) @@ -139,7 +144,7 @@ async def test_msat_wallet_surplus_is_not_rounded( """Either bound can bind, and neither is rounded to whole sats.""" get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0}) with ( - _liabilities({}, total_sats=0), + _env(), patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), patch( diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index 9e57f901..f1315835 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -14,7 +14,7 @@ from collections.abc import Callable, Coroutine from contextlib import asynccontextmanager from pathlib import Path from typing import Any -from unittest.mock import ANY, AsyncMock, MagicMock, call, patch +from unittest.mock import ANY, AsyncMock, MagicMock, patch import pytest @@ -248,11 +248,13 @@ async def test_periodic_payout_isolates_failing_mint() -> None: # The bad mint raised on get_wallet for both units, yet the good mint was # still reached and paid out for both units — failures are isolated. - for unit in ("sat", "msat"): - assert ( - call("http://good:3338", unit, force_reload_proofs=True) - in get_wallet.await_args_list - ) + good_reloads = [ + c + for c in get_wallet.await_args_list + if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs") + ] + # Two payout reads, plus two cross-wallet reads for the global payout bound. + assert len(good_reloads) == 4 assert raw_send.await_count == 2 # good mint paid for both units diff --git a/tests/unit/test_user_liability_for_mint_and_unit.py b/tests/unit/test_user_liability_for_mint_and_unit.py index 4ce00941..7d8abbc1 100644 --- a/tests/unit/test_user_liability_for_mint_and_unit.py +++ b/tests/unit/test_user_liability_for_mint_and_unit.py @@ -1,9 +1,4 @@ -"""Real-DB coverage for db.user_liability_for_mint_and_unit. - -Verifies the per-mint liability query that bounds owner payout: it sums key -balances and unresolved refund claims for one (mint_url, unit), excludes -resolved claims and other mints/units, and drops keys with no refund mint. -""" +"""Real-DB coverage for the per-mint liability query that bounds owner payout.""" from typing import AsyncGenerator @@ -88,7 +83,7 @@ async def test_sums_key_balances_for_the_mint_and_unit(session: AsyncSession) -> @pytest.mark.asyncio async def test_adds_unresolved_refunds_to_key_balances(session: AsyncSession) -> None: - # One open claim per key, so each unresolved status needs its own key. + # Only one pending/ambiguous claim per key is allowed. await _add_key(session, "a", 1000) await _add_refund(session, "a", 300, "pending") await _add_key(session, "b", 0) From 92fbc2e1ec85ec48625bd1e6ced9e5a7b211c353 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 24 Sep 2026 18:27:27 +0200 Subject: [PATCH 4/4] clean up --- routstr/wallet.py | 88 +++++++++++++------- tests/unit/test_lnurl_change.py | 12 ++- tests/unit/test_payout_liability_bounds.py | 94 +++++++++++++--------- tests/unit/test_periodic_payout.py | 20 +++-- tests/unit/test_wallet.py | 53 ++++++++++++ 5 files changed, 190 insertions(+), 77 deletions(-) diff --git a/routstr/wallet.py b/routstr/wallet.py index c69f1aaa..1f2a9e0e 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -15,6 +15,7 @@ import httpx from cashu.core.base import MeltQuote, Proof, Token from cashu.core.mint_info import MintInfo as _CashuMintInfo from cashu.wallet.crud import get_keysets as get_cashu_keysets +from cashu.wallet.crud import get_proofs as get_cashu_proofs from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet as _CashuWallet from pydantic_core import PydanticUndefined @@ -184,9 +185,12 @@ class Wallet(_CashuWallet): pass await self.load_mint_keysets(force_old_keysets) - await self.activate_keyset(keyset_id) await self.load_mint_info(reload=True) + # Arm on the fetch, not the activation: a unit the mint does not + # serve makes ``activate_keyset`` raise, and arming after it would + # refetch keysets on every call. _mint_metadata_last_load[mint_url] = time.monotonic() + await self.activate_keyset(keyset_id) class MintConnectionError(Exception): @@ -708,27 +712,30 @@ def _to_msats(amount: int, unit: str) -> int: async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int: """Sum unreserved proofs of every other trusted wallet, in msats. - This total only ever raises the payout ceiling, so proofs are reloaded from - the local db: a cached snapshot up to 30s stale could still hide another - process's reservation. A wallet that cannot be loaded counts as empty, - which can only shrink the owner surplus. + Every wallet shares one db, so two queries answer for all of them. Loading + a wallet per mint and unit instead refetched keysets from each mint on + every call and rate-limited them. + + Read fresh, not from a wallet's snapshot: this total only ever raises the + payout ceiling, and a snapshot up to 30s stale could hide another + process's reservation. """ + wallet = await get_wallet(mint_url, unit, load=False) + trusted = set(_mints_to_inspect()) + origins: dict[str, tuple[str, str]] = {} + for keyset in await get_cashu_keysets(db=wallet.db): + keyset_unit = keyset.unit if isinstance(keyset.unit, str) else keyset.unit.name + origin = (keyset.mint_url, keyset_unit) + if origin == (mint_url, unit): + continue + if keyset.mint_url in trusted and keyset_unit in ("sat", "msat"): + origins[keyset.id] = origin total = 0 - for other_mint in _mints_to_inspect(): - for other_unit in ("sat", "msat"): - if (other_mint, other_unit) == (mint_url, unit): - continue - try: - wallet = await get_wallet( - other_mint, other_unit, force_reload_proofs=True - ) - except Exception as e: - logger.debug(f"Wallet {other_mint} {other_unit} excluded: {e}") - continue - proofs = get_proofs_per_mint_and_unit( - wallet, other_mint, other_unit, not_reserved=True - ) - total += _to_msats(sum(proof.amount for proof in proofs), other_unit) + for proof in await get_cashu_proofs(db=wallet.db): + proof_origin = origins.get(proof.id) + if proof_origin is None or proof.reserved: + continue + total += _to_msats(proof.amount, proof_origin[1]) return total @@ -1203,6 +1210,9 @@ _wallets: dict[str, Wallet] = {} # Proofs require a shorter refresh interval than remote mint metadata. _wallet_last_load: dict[str, float] = {} _wallet_last_mint_load: dict[str, float] = {} +# Metadata loads the mint answered but that left the wallet unusable, replayed +# for the reload interval so the failure costs one request, not one per call. +_wallet_mint_load_errors: dict[str, tuple[float, Exception]] = {} _wallet_load_locks: dict[str, asyncio.Lock] = {} @@ -1230,16 +1240,34 @@ async def get_wallet( or last_mint_load is None or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS ): - await run_mint_operation( - lambda: ( - _wallets[id].load_mint(force_refresh=True) - if force_reload - else _wallets[id].load_mint() - ), - op_name="load_mint", - mint_url=mint_url, - retry_on_rate_limit=retry_on_rate_limit, - ) + cached_error = _wallet_mint_load_errors.get(id) + if ( + not force_reload + and cached_error is not None + and now - cached_error[0] < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS + ): + raise cached_error[1] + try: + await run_mint_operation( + lambda: ( + _wallets[id].load_mint(force_refresh=True) + if force_reload + else _wallets[id].load_mint() + ), + op_name="load_mint", + mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, + ) + except Exception as error: + # Transport failures and 429s stay retryable; the rate + # guard owns those. Anything else means the mint answered + # and still cannot serve this wallet. + if not ( + is_mint_connection_error(error) or _is_mint_rate_limited(error) + ): + _wallet_mint_load_errors[id] = (time.monotonic(), error) + raise + _wallet_mint_load_errors.pop(id, None) _wallet_last_mint_load[id] = time.monotonic() if load_proofs: diff --git a/tests/unit/test_lnurl_change.py b/tests/unit/test_lnurl_change.py index fc62d02c..90ea1a9b 100644 --- a/tests/unit/test_lnurl_change.py +++ b/tests/unit/test_lnurl_change.py @@ -1,4 +1,4 @@ -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Iterator from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch @@ -15,6 +15,16 @@ from routstr.mint import MintRateGuard from routstr.wallet import _payout_mint_and_unit +@pytest.fixture(autouse=True) +def empty_cross_wallet_proofs() -> Iterator[None]: + """No other wallet holds proofs, so only this wallet's own bound applies.""" + with ( + patch("routstr.wallet.get_cashu_keysets", AsyncMock(return_value=[])), + patch("routstr.wallet.get_cashu_proofs", AsyncMock(return_value=[])), + ): + yield + + @pytest.mark.asyncio @pytest.mark.parametrize("unit,scale", [("sat", 1), ("msat", 1000)]) @pytest.mark.parametrize( diff --git a/tests/unit/test_payout_liability_bounds.py b/tests/unit/test_payout_liability_bounds.py index 98979e99..2fb8ae0d 100644 --- a/tests/unit/test_payout_liability_bounds.py +++ b/tests/unit/test_payout_liability_bounds.py @@ -23,20 +23,34 @@ async def _session() -> AsyncIterator[Mock]: yield Mock() -def _wallets( - sat_proofs: dict[str, int], unreachable: frozenset[str] = frozenset() -) -> tuple[AsyncMock, Mock]: - """Fake get_wallet/get_proofs for sat wallets; msat wallets are unsupported.""" +def _keyset_id(mint_url: str, unit: str) -> str: + return f"{mint_url}|{unit}" - async def get_wallet(mint_url: str, unit: str, **_: object) -> Mock: - if unit != "sat" or mint_url in unreachable: - raise ValueError("unsupported") - return Mock(url=mint_url) - def get_proofs(wallet: Mock, mint_url: str, unit: str, **_: object) -> list[Mock]: - return [Mock(amount=sat_proofs[mint_url])] - - return AsyncMock(side_effect=get_wallet), Mock(side_effect=get_proofs) +@contextmanager +def _wallet_db( + sat_proofs: dict[str, int], reserved: frozenset[str] = frozenset() +) -> Iterator[AsyncMock]: + """One sat keyset per mint, one proof behind it. Yields the get_wallet mock.""" + keysets = [ + Mock(id=_keyset_id(mint_url, "sat"), mint_url=mint_url, unit="sat") + for mint_url in sat_proofs + ] + proofs = [ + Mock( + id=_keyset_id(mint_url, "sat"), + amount=amount, + reserved=mint_url in reserved, + ) + for mint_url, amount in sat_proofs.items() + ] + get_wallet = AsyncMock(return_value=Mock(url=MINT_B, db=Mock())) + with ( + patch("routstr.wallet.get_wallet", get_wallet), + patch("routstr.wallet.get_cashu_keysets", AsyncMock(return_value=keysets)), + patch("routstr.wallet.get_cashu_proofs", AsyncMock(return_value=proofs)), + ): + yield get_wallet @contextmanager @@ -70,11 +84,9 @@ def _liabilities(per_mint_sats: dict[str, int], total_sats: int) -> Iterator[Non @pytest.mark.asyncio async def test_owner_balance_keeps_only_the_wallets_own_liability() -> None: - get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) with ( _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + _wallet_db({MINT_A: 400, MINT_B: 270}), ): assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 236 assert await _owner_balance_for_mint_and_unit(MINT_A, "sat", 400) == 184 @@ -83,25 +95,31 @@ async def test_owner_balance_keeps_only_the_wallets_own_liability() -> None: @pytest.mark.asyncio async def test_owner_balance_never_exceeds_global_surplus() -> None: """Liability nobody declared against a mint is still covered in aggregate.""" - get_wallet, get_proofs = _wallets({MINT_A: 100, MINT_B: 270}) with ( _liabilities({}, total_sats=250), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + _wallet_db({MINT_A: 100, MINT_B: 270}), ): assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 120 @pytest.mark.asyncio -async def test_unloadable_wallet_counts_as_empty() -> None: - """Shrinks the surplus rather than inflating it.""" - get_wallet, get_proofs = _wallets( - {MINT_A: 400, MINT_B: 270}, unreachable=frozenset({MINT_A}) - ) +async def test_proofs_of_an_untrusted_mint_do_not_raise_the_bound() -> None: + """Only configured mints back the global surplus.""" with ( _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + patch.object(settings, "cashu_mints", [MINT_B]), + patch.object(settings, "primary_mint", MINT_B), + _wallet_db({MINT_A: 400, MINT_B: 270}), + ): + assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 20 + + +@pytest.mark.asyncio +async def test_reserved_proofs_do_not_raise_the_bound() -> None: + """Another process may already be spending them.""" + with ( + _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), + _wallet_db({MINT_A: 400, MINT_B: 270}, reserved=frozenset({MINT_A})), ): assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 20 @@ -109,28 +127,24 @@ async def test_unloadable_wallet_counts_as_empty() -> None: @pytest.mark.asyncio async def test_duplicate_configured_mint_is_counted_once() -> None: """A mint listed twice in CASHU_MINTS would otherwise raise the global bound.""" - get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) with ( _liabilities({}, total_sats=600), patch.object(settings, "cashu_mints", [MINT_A, MINT_A, MINT_B]), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + _wallet_db({MINT_A: 400, MINT_B: 270}), ): assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 70 @pytest.mark.asyncio -async def test_other_wallets_are_read_from_fresh_local_proofs() -> None: - """A stale snapshot of another wallet would raise the global bound.""" - get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) +async def test_cross_wallet_bound_asks_no_mint_for_metadata() -> None: + """The sum is local. Loading a wallet per mint and unit rate-limited mints.""" with ( _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + _wallet_db({MINT_A: 400, MINT_B: 270}) as get_wallet, ): await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) assert get_wallet.await_args_list - assert all(c.kwargs.get("force_reload_proofs") for c in get_wallet.await_args_list) + assert all(c.kwargs.get("load") is False for c in get_wallet.await_args_list) @pytest.mark.asyncio @@ -142,11 +156,9 @@ async def test_msat_wallet_surplus_is_not_rounded( mint_liability: int, total_liability: int, expected: int ) -> None: """Either bound can bind, and neither is rounded to whole sats.""" - get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0}) with ( _env(), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + _wallet_db({}), patch( "routstr.wallet.db.user_liability_for_mint_and_unit", AsyncMock(return_value=mint_liability), @@ -161,14 +173,16 @@ async def test_msat_wallet_surplus_is_not_rounded( @pytest.mark.asyncio async def test_payout_sends_the_smaller_wallets_surplus() -> None: - get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270}) send = AsyncMock(return_value=236_000) with ( _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), patch.object(settings, "min_payout_sat", 50), patch.object(settings, "max_payout_sat", 250_000), - patch("routstr.wallet.get_wallet", get_wallet), - patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), + _wallet_db({MINT_A: 400, MINT_B: 270}), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + Mock(return_value=[Mock(amount=270)]), + ), patch( "routstr.wallet.slow_filter_spend_proofs", AsyncMock(side_effect=lambda proofs, wallet: proofs), diff --git a/tests/unit/test_periodic_payout.py b/tests/unit/test_periodic_payout.py index f1315835..4286bb40 100644 --- a/tests/unit/test_periodic_payout.py +++ b/tests/unit/test_periodic_payout.py @@ -10,7 +10,7 @@ Covers two regressions from the auto-payout / primary-mint audit mint/units in the same cycle (the try/except is now per mint/unit). """ -from collections.abc import Callable, Coroutine +from collections.abc import Callable, Coroutine, Iterator from contextlib import asynccontextmanager from pathlib import Path from typing import Any @@ -26,6 +26,16 @@ from routstr.wallet import ( ) +@pytest.fixture(autouse=True) +def empty_cross_wallet_proofs() -> Iterator[None]: + """No other wallet holds proofs, so only this wallet's own bound applies.""" + with ( + patch("routstr.wallet.get_cashu_keysets", AsyncMock(return_value=[])), + patch("routstr.wallet.get_cashu_proofs", AsyncMock(return_value=[])), + ): + yield + + @pytest.fixture(autouse=True) def isolate_wallet_lock(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( @@ -202,9 +212,7 @@ async def test_periodic_payout_isolates_failing_mint() -> None: """A failing mint does not prevent payout for the other mints.""" from routstr.core.settings import settings - async def _get_wallet( - mint_url: str, unit: str, force_reload_proofs: bool = False - ) -> MagicMock: + async def _get_wallet(mint_url: str, unit: str, **_: object) -> MagicMock: if mint_url == "http://bad:3338": raise RuntimeError("mint unreachable") return MagicMock() @@ -253,8 +261,8 @@ async def test_periodic_payout_isolates_failing_mint() -> None: for c in get_wallet.await_args_list if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs") ] - # Two payout reads, plus two cross-wallet reads for the global payout bound. - assert len(good_reloads) == 4 + # One proof read per unit, and no extra mint load for the bound. + assert len(good_reloads) == 2 assert raw_send.await_count == 2 # good mint paid for both units diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 8054f91c..30b71f09 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -48,6 +48,7 @@ def isolate_wallet_runtime_state( wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() wallet_module._wallet_last_mint_load.clear() + wallet_module._wallet_mint_load_errors.clear() wallet_module._wallet_load_locks.clear() wallet_module._mint_metadata_last_load.clear() wallet_module._mint_metadata_load_locks.clear() @@ -57,6 +58,7 @@ def isolate_wallet_runtime_state( wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() wallet_module._wallet_last_mint_load.clear() + wallet_module._wallet_mint_load_errors.clear() wallet_module._wallet_load_locks.clear() wallet_module._mint_metadata_last_load.clear() wallet_module._mint_metadata_load_locks.clear() @@ -149,6 +151,57 @@ async def test_get_wallet_force_reload_bypasses_reload_interval() -> None: assert mock_wallet.load_proofs.await_count == 2 +@pytest.mark.asyncio +async def test_unservable_mint_load_is_not_retried_every_call() -> None: + """Retrying it per call refetched keysets and got the node rate-limited.""" + from routstr.wallet import get_wallet + + failure = Exception("No active keyset found for unit msat.") + mock_wallet = Mock( + load_mint=AsyncMock(side_effect=failure), load_proofs=AsyncMock() + ) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + for _ in range(3): + with pytest.raises(Exception, match="No active keyset"): + await get_wallet("http://mint:3338", "msat") + + assert mock_wallet.load_mint.await_count == 1 + + +@pytest.mark.asyncio +async def test_unreachable_mint_load_stays_retryable() -> None: + """Transport failures are the rate guard's job, not the metadata throttle's.""" + from routstr.wallet import get_wallet + + failure = httpx.ConnectError("mint unreachable") + mock_wallet = Mock( + load_mint=AsyncMock(side_effect=failure), load_proofs=AsyncMock() + ) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + for _ in range(2): + with pytest.raises(Exception): + await get_wallet("http://mint:3338", "sat") + + assert mock_wallet.load_mint.await_count == 2 + + +@pytest.mark.asyncio +async def test_force_reload_retries_an_unservable_mint_load() -> None: + from routstr.wallet import get_wallet + + failure = Exception("No active keyset found for unit msat.") + mock_wallet = Mock( + load_mint=AsyncMock(side_effect=failure), load_proofs=AsyncMock() + ) + with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)): + with pytest.raises(Exception, match="No active keyset"): + await get_wallet("http://mint:3338", "msat") + with pytest.raises(Exception, match="No active keyset"): + await get_wallet("http://mint:3338", "msat", force_reload=True) + + assert mock_wallet.load_mint.await_count == 2 + + @pytest.mark.asyncio async def test_get_wallet_force_reload_proofs_keeps_cached_keysets() -> None: from routstr.wallet import get_wallet