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=[]),