This commit is contained in:
9qeklajc
2026-09-24 18:27:27 +02:00
parent 90be8e6792
commit 92fbc2e1ec
5 changed files with 190 additions and 77 deletions
+58 -30
View File
@@ -15,6 +15,7 @@ import httpx
from cashu.core.base import MeltQuote, Proof, Token from cashu.core.base import MeltQuote, Proof, Token
from cashu.core.mint_info import MintInfo as _CashuMintInfo 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_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.helpers import deserialize_token_from_string
from cashu.wallet.wallet import Wallet as _CashuWallet from cashu.wallet.wallet import Wallet as _CashuWallet
from pydantic_core import PydanticUndefined from pydantic_core import PydanticUndefined
@@ -184,9 +185,12 @@ class Wallet(_CashuWallet):
pass pass
await self.load_mint_keysets(force_old_keysets) await self.load_mint_keysets(force_old_keysets)
await self.activate_keyset(keyset_id)
await self.load_mint_info(reload=True) 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() _mint_metadata_last_load[mint_url] = time.monotonic()
await self.activate_keyset(keyset_id)
class MintConnectionError(Exception): 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: async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int:
"""Sum unreserved proofs of every other trusted wallet, in msats. """Sum unreserved proofs of every other trusted wallet, in msats.
This total only ever raises the payout ceiling, so proofs are reloaded from Every wallet shares one db, so two queries answer for all of them. Loading
the local db: a cached snapshot up to 30s stale could still hide another a wallet per mint and unit instead refetched keysets from each mint on
process's reservation. A wallet that cannot be loaded counts as empty, every call and rate-limited them.
which can only shrink the owner surplus.
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 total = 0
for other_mint in _mints_to_inspect(): for proof in await get_cashu_proofs(db=wallet.db):
for other_unit in ("sat", "msat"): proof_origin = origins.get(proof.id)
if (other_mint, other_unit) == (mint_url, unit): if proof_origin is None or proof.reserved:
continue continue
try: total += _to_msats(proof.amount, proof_origin[1])
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)
return total return total
@@ -1203,6 +1210,9 @@ _wallets: dict[str, Wallet] = {}
# Proofs require a shorter refresh interval than remote mint metadata. # Proofs require a shorter refresh interval than remote mint metadata.
_wallet_last_load: dict[str, float] = {} _wallet_last_load: dict[str, float] = {}
_wallet_last_mint_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] = {} _wallet_load_locks: dict[str, asyncio.Lock] = {}
@@ -1230,16 +1240,34 @@ async def get_wallet(
or last_mint_load is None or last_mint_load is None
or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
): ):
await run_mint_operation( cached_error = _wallet_mint_load_errors.get(id)
lambda: ( if (
_wallets[id].load_mint(force_refresh=True) not force_reload
if force_reload and cached_error is not None
else _wallets[id].load_mint() and now - cached_error[0] < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
), ):
op_name="load_mint", raise cached_error[1]
mint_url=mint_url, try:
retry_on_rate_limit=retry_on_rate_limit, 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() _wallet_last_mint_load[id] = time.monotonic()
if load_proofs: if load_proofs:
+11 -1
View File
@@ -1,4 +1,4 @@
from collections.abc import AsyncIterator from collections.abc import AsyncIterator, Iterator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
@@ -15,6 +15,16 @@ from routstr.mint import MintRateGuard
from routstr.wallet import _payout_mint_and_unit 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.asyncio
@pytest.mark.parametrize("unit,scale", [("sat", 1), ("msat", 1000)]) @pytest.mark.parametrize("unit,scale", [("sat", 1), ("msat", 1000)])
@pytest.mark.parametrize( @pytest.mark.parametrize(
+54 -40
View File
@@ -23,20 +23,34 @@ async def _session() -> AsyncIterator[Mock]:
yield Mock() yield Mock()
def _wallets( def _keyset_id(mint_url: str, unit: str) -> str:
sat_proofs: dict[str, int], unreachable: frozenset[str] = frozenset() return f"{mint_url}|{unit}"
) -> 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]: @contextmanager
return [Mock(amount=sat_proofs[mint_url])] def _wallet_db(
sat_proofs: dict[str, int], reserved: frozenset[str] = frozenset()
return AsyncMock(side_effect=get_wallet), Mock(side_effect=get_proofs) ) -> 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 @contextmanager
@@ -70,11 +84,9 @@ def _liabilities(per_mint_sats: dict[str, int], total_sats: int) -> Iterator[Non
@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:
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),
patch("routstr.wallet.get_wallet", get_wallet), _wallet_db({MINT_A: 400, MINT_B: 270}),
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_B, "sat", 270) == 236
assert await _owner_balance_for_mint_and_unit(MINT_A, "sat", 400) == 184 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 @pytest.mark.asyncio
async def test_owner_balance_never_exceeds_global_surplus() -> None: async def test_owner_balance_never_exceeds_global_surplus() -> None:
"""Liability nobody declared against a mint is still covered in aggregate.""" """Liability nobody declared against a mint is still covered in aggregate."""
get_wallet, get_proofs = _wallets({MINT_A: 100, MINT_B: 270})
with ( with (
_liabilities({}, total_sats=250), _liabilities({}, total_sats=250),
patch("routstr.wallet.get_wallet", get_wallet), _wallet_db({MINT_A: 100, MINT_B: 270}),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs),
): ):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 120 assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 120
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_unloadable_wallet_counts_as_empty() -> None: async def test_proofs_of_an_untrusted_mint_do_not_raise_the_bound() -> None:
"""Shrinks the surplus rather than inflating it.""" """Only configured mints back the global surplus."""
get_wallet, get_proofs = _wallets(
{MINT_A: 400, MINT_B: 270}, unreachable=frozenset({MINT_A})
)
with ( with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
patch("routstr.wallet.get_wallet", get_wallet), patch.object(settings, "cashu_mints", [MINT_B]),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), 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 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 @pytest.mark.asyncio
async def test_duplicate_configured_mint_is_counted_once() -> None: async def test_duplicate_configured_mint_is_counted_once() -> None:
"""A mint listed twice in CASHU_MINTS would otherwise raise the global bound.""" """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 ( with (
_liabilities({}, total_sats=600), _liabilities({}, total_sats=600),
patch.object(settings, "cashu_mints", [MINT_A, MINT_A, MINT_B]), patch.object(settings, "cashu_mints", [MINT_A, MINT_A, MINT_B]),
patch("routstr.wallet.get_wallet", get_wallet), _wallet_db({MINT_A: 400, MINT_B: 270}),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs),
): ):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 70 assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 70
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_other_wallets_are_read_from_fresh_local_proofs() -> None: async def test_cross_wallet_bound_asks_no_mint_for_metadata() -> None:
"""A stale snapshot of another wallet would raise the global bound.""" """The sum is local. Loading a wallet per mint and unit rate-limited mints."""
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),
patch("routstr.wallet.get_wallet", get_wallet), _wallet_db({MINT_A: 400, MINT_B: 270}) as get_wallet,
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs),
): ):
await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270)
assert get_wallet.await_args_list 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 @pytest.mark.asyncio
@@ -142,11 +156,9 @@ async def test_msat_wallet_surplus_is_not_rounded(
mint_liability: int, total_liability: int, expected: int mint_liability: int, total_liability: int, expected: int
) -> None: ) -> None:
"""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})
with ( with (
_env(), _env(),
patch("routstr.wallet.get_wallet", get_wallet), _wallet_db({}),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs),
patch( patch(
"routstr.wallet.db.user_liability_for_mint_and_unit", "routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=mint_liability), AsyncMock(return_value=mint_liability),
@@ -161,14 +173,16 @@ async def test_msat_wallet_surplus_is_not_rounded(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_payout_sends_the_smaller_wallets_surplus() -> None: 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) send = AsyncMock(return_value=236_000)
with ( with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250), _liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
patch.object(settings, "min_payout_sat", 50), patch.object(settings, "min_payout_sat", 50),
patch.object(settings, "max_payout_sat", 250_000), patch.object(settings, "max_payout_sat", 250_000),
patch("routstr.wallet.get_wallet", get_wallet), _wallet_db({MINT_A: 400, MINT_B: 270}),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs), patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
Mock(return_value=[Mock(amount=270)]),
),
patch( patch(
"routstr.wallet.slow_filter_spend_proofs", "routstr.wallet.slow_filter_spend_proofs",
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
+14 -6
View File
@@ -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). 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 contextlib import asynccontextmanager
from pathlib import Path from pathlib import Path
from typing import Any 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) @pytest.fixture(autouse=True)
def isolate_wallet_lock(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: def isolate_wallet_lock(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr( 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.""" """A failing mint does not prevent payout for the other mints."""
from routstr.core.settings import settings from routstr.core.settings import settings
async def _get_wallet( async def _get_wallet(mint_url: str, unit: str, **_: object) -> MagicMock:
mint_url: str, unit: str, force_reload_proofs: bool = False
) -> MagicMock:
if mint_url == "http://bad:3338": if mint_url == "http://bad:3338":
raise RuntimeError("mint unreachable") raise RuntimeError("mint unreachable")
return MagicMock() return MagicMock()
@@ -253,8 +261,8 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
for c in get_wallet.await_args_list for c in get_wallet.await_args_list
if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs") 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. # One proof read per unit, and no extra mint load for the bound.
assert len(good_reloads) == 4 assert len(good_reloads) == 2
assert raw_send.await_count == 2 # good mint paid for both units assert raw_send.await_count == 2 # good mint paid for both units
+53
View File
@@ -48,6 +48,7 @@ def isolate_wallet_runtime_state(
wallet_module._wallets.clear() wallet_module._wallets.clear()
wallet_module._wallet_last_load.clear() wallet_module._wallet_last_load.clear()
wallet_module._wallet_last_mint_load.clear() wallet_module._wallet_last_mint_load.clear()
wallet_module._wallet_mint_load_errors.clear()
wallet_module._wallet_load_locks.clear() wallet_module._wallet_load_locks.clear()
wallet_module._mint_metadata_last_load.clear() wallet_module._mint_metadata_last_load.clear()
wallet_module._mint_metadata_load_locks.clear() wallet_module._mint_metadata_load_locks.clear()
@@ -57,6 +58,7 @@ def isolate_wallet_runtime_state(
wallet_module._wallets.clear() wallet_module._wallets.clear()
wallet_module._wallet_last_load.clear() wallet_module._wallet_last_load.clear()
wallet_module._wallet_last_mint_load.clear() wallet_module._wallet_last_mint_load.clear()
wallet_module._wallet_mint_load_errors.clear()
wallet_module._wallet_load_locks.clear() wallet_module._wallet_load_locks.clear()
wallet_module._mint_metadata_last_load.clear() wallet_module._mint_metadata_last_load.clear()
wallet_module._mint_metadata_load_locks.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 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 @pytest.mark.asyncio
async def test_get_wallet_force_reload_proofs_keeps_cached_keysets() -> None: async def test_get_wallet_force_reload_proofs_keeps_cached_keysets() -> None:
from routstr.wallet import get_wallet from routstr.wallet import get_wallet