From 92fbc2e1ec85ec48625bd1e6ced9e5a7b211c353 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 24 Sep 2026 18:27:27 +0200 Subject: [PATCH] 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