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")