refactor: simplify payout liability tests and tighten reload assertion

This commit is contained in:
9qeklajc
2026-09-23 19:55:26 +02:00
parent b663b1d83a
commit 90be8e6792
4 changed files with 38 additions and 37 deletions
+4 -5
View File
@@ -830,9 +830,7 @@ async def _prepare_bolt11_payment(invoice: str) -> Bolt11PaymentPlan:
) )
if owner_balance < required: if owner_balance < required:
continue continue
owner_balance_msats = ( owner_balance_msats = _to_msats(owner_balance, unit)
owner_balance * 1000 if unit == "sat" else owner_balance
)
candidates.append( candidates.append(
(owner_balance_msats, wallet, proofs, quote, mint_url, unit) (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 return
# Read liabilities after the proofs snapshot and settle delay while the # Read liabilities and the other wallets' proofs after this wallet's proofs
# wallet operation guard excludes concurrent proof mutation and crediting. # snapshot and settle delay, while the wallet operation guard excludes
# concurrent proof mutation and crediting.
try: try:
available_balance = await _owner_balance_for_mint_and_unit( available_balance = await _owner_balance_for_mint_and_unit(
mint_url, unit, sum(proof.amount for proof in proofs) mint_url, unit, sum(proof.amount for proof in proofs)
+24 -19
View File
@@ -6,7 +6,7 @@ than the whole liability, so only the largest wallet could ever pay out.
""" """
from collections.abc import AsyncIterator, Iterator from collections.abc import AsyncIterator, Iterator
from contextlib import ExitStack, asynccontextmanager, contextmanager from contextlib import asynccontextmanager, contextmanager
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
import pytest import pytest
@@ -39,32 +39,37 @@ def _wallets(
return AsyncMock(side_effect=get_wallet), Mock(side_effect=get_proofs) 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 @contextmanager
def _liabilities(per_mint_sats: dict[str, int], total_sats: int) -> Iterator[None]: 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: async def per_mint(_session: object, mint_url: str, unit: str) -> int:
return per_mint_sats.get(mint_url, 0) * 1000 return per_mint_sats.get(mint_url, 0) * 1000
with ExitStack() as stack: with (
for target in ( _env(),
patch( patch(
"routstr.wallet.db.user_liability_for_mint_and_unit", "routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(side_effect=per_mint), AsyncMock(side_effect=per_mint),
), ),
patch( patch(
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=total_sats * 1000), 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 yield
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_owner_balance_keeps_only_the_wallets_own_liability() -> None: 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}) get_wallet, get_proofs = _wallets({MINT_A: 400, MINT_B: 270})
with ( with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), _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 @pytest.mark.asyncio
async def test_unloadable_wallet_counts_as_empty() -> None: 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( get_wallet, get_proofs = _wallets(
{MINT_A: 400, MINT_B: 270}, unreachable=frozenset({MINT_A}) {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.""" """Either bound can bind, and neither is rounded to whole sats."""
get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0}) get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0})
with ( with (
_liabilities({}, total_sats=0), _env(),
patch("routstr.wallet.get_wallet", get_wallet), patch("routstr.wallet.get_wallet", get_wallet),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs),
patch( patch(
+8 -6
View File
@@ -14,7 +14,7 @@ from collections.abc import Callable, Coroutine
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch from unittest.mock import ANY, AsyncMock, MagicMock, patch
import pytest 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 # 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. # still reached and paid out for both units — failures are isolated.
for unit in ("sat", "msat"): good_reloads = [
assert ( c
call("http://good:3338", unit, force_reload_proofs=True) for c in get_wallet.await_args_list
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 assert raw_send.await_count == 2 # good mint paid for both units
@@ -1,9 +1,4 @@
"""Real-DB coverage for db.user_liability_for_mint_and_unit. """Real-DB coverage for the per-mint liability query that bounds owner payout."""
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 from typing import AsyncGenerator
@@ -88,7 +83,7 @@ async def test_sums_key_balances_for_the_mint_and_unit(session: AsyncSession) ->
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_adds_unresolved_refunds_to_key_balances(session: AsyncSession) -> None: 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_key(session, "a", 1000)
await _add_refund(session, "a", 300, "pending") await _add_refund(session, "a", 300, "pending")
await _add_key(session, "b", 0) await _add_key(session, "b", 0)