From 90be8e6792790afd9b66d47a676907287b4ebd3b Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 23 Sep 2026 19:55:26 +0200 Subject: [PATCH] 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)