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
+47 -19
View File
@@ -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):
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
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)
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,6 +1240,14 @@ async def get_wallet(
or last_mint_load is None
or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
):
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)
@@ -1240,6 +1258,16 @@ async def get_wallet(
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:
+11 -1
View File
@@ -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(
+54 -40
View File
@@ -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),
+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).
"""
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
+53
View File
@@ -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