Merge pull request #767 from Routstr/fix/payout-liability-per-mint

fix: bound owner payout by per-mint liability instead of total
This commit is contained in:
9qeklajc
2026-09-24 22:31:38 +02:00
committed by GitHub
8 changed files with 584 additions and 46 deletions
+28
View File
@@ -1098,6 +1098,34 @@ async def total_user_liability(db_session: AsyncSession) -> int:
return int(result.one() or 0) return int(result.one() or 0)
async def user_liability_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str
) -> int:
"""Return outstanding user funds that refund from one mint and unit, in msats.
Single statement, for the same atomicity reason as ``total_user_liability``.
"""
key_balances = (
select(func.coalesce(func.sum(ApiKey.balance), 0))
.where(
col(ApiKey.refund_mint_url) == mint_url,
col(ApiKey.refund_currency) == unit,
)
.scalar_subquery()
)
unresolved_refunds = (
select(func.coalesce(func.sum(Refund.amount_msats), 0))
.where(
col(Refund.status).in_(REFUND_UNRESOLVED_STATUSES),
col(Refund.mint_url) == mint_url,
col(Refund.unit) == unit,
)
.scalar_subquery()
)
result = await db_session.exec(select(key_balances + unresolved_refunds))
return int(result.one() or 0)
async def balance_for_mint_and_unit( async def balance_for_mint_and_unit(
db_session: AsyncSession, mint_url: str, unit: str db_session: AsyncSession, mint_url: str, unit: str
) -> int: ) -> int:
+98 -35
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
@@ -121,7 +122,7 @@ def _msats_to_sats_ceil(amount: int) -> int:
def _mints_to_inspect() -> list[str]: def _mints_to_inspect() -> list[str]:
"""Return configured mints plus the primary mint, without duplicates.""" """Return configured mints plus the primary mint, without duplicates."""
mint_urls = list(settings.cashu_mints) mint_urls = list(dict.fromkeys(settings.cashu_mints))
if settings.primary_mint and settings.primary_mint not in mint_urls: if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint) mint_urls.append(settings.primary_mint)
return mint_urls return mint_urls
@@ -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):
@@ -701,18 +705,64 @@ class Bolt11PaymentPlan:
return maximum if self.unit == "sat" else (maximum + 999) // 1000 return maximum if self.unit == "sat" else (maximum + 999) // 1000
def _to_msats(amount: int, unit: str) -> int:
return _sats_to_msats(amount) if unit == "sat" else amount
async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int:
"""Sum unreserved proofs of every other trusted wallet, in msats.
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 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
async def _owner_balance_for_mint_and_unit( async def _owner_balance_for_mint_and_unit(
mint_url: str, unit: str, proofs_balance: int mint_url: str, unit: str, proofs_balance: int
) -> int: ) -> int:
"""Return spendable node-owned funds without crossing user liabilities.""" """Return owner funds in one wallet, in that wallet's unit.
A key's refund mint is a preference, not funding provenance: a key topped
up from a second mint keeps its original refund mint. Hence two bounds —
the per-mint one keeps refunds serviceable from the mint they name, the
global one stops misattributed customer funds being paid out as profit.
"""
others_msats = await _other_wallets_unreserved_msats(mint_url, unit)
async with db.create_session() as session: async with db.create_session() as session:
# Refund mint is a preference, not funding provenance. Mirror payout's mint_liability = await db.user_liability_for_mint_and_unit(
# conservative rule and protect the full liability at every mint. session, mint_url, unit
user_liability = await db.total_user_liability(session) )
# API-key balances are stored in msats. Cashu ``sat`` proofs are not. total_liability = await db.total_user_liability(session)
if unit == "sat": proofs_msats = _to_msats(proofs_balance, unit)
user_liability = _msats_to_sats_ceil(user_liability) surplus_msats = min(
return max(0, proofs_balance - user_liability) proofs_msats - mint_liability,
proofs_msats + others_msats - total_liability,
)
# Cashu ``sat`` proofs are whole sats.
surplus = _msats_to_sats(surplus_msats) if unit == "sat" else surplus_msats
return max(0, surplus)
async def maximum_owner_cashu_balance_sats() -> int: async def maximum_owner_cashu_balance_sats() -> int:
@@ -787,9 +837,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)
) )
@@ -1162,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] = {}
@@ -1189,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:
@@ -1650,15 +1719,13 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
) )
return return
# Fetch liability 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:
async with db.create_session() as session: available_balance = await _owner_balance_for_mint_and_unit(
# ApiKey stores a refund preference, not funding provenance. Until mint_url, unit, sum(proof.amount for proof in proofs)
# liabilities have a durable per-credit ledger, subtract the total )
# liability from every wallet rather than risk calling customer
# funds owner profit on the wrong mint.
user_balance = await db.total_user_liability(session)
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Error in periodic payout cycle: {type(e).__name__}", f"Error in periodic payout cycle: {type(e).__name__}",
@@ -1667,10 +1734,6 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
return return
try: try:
if unit == "sat":
user_balance = _msats_to_sats_ceil(user_balance)
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
max_amount = ( max_amount = (
settings.max_payout_sat settings.max_payout_sat
if unit == "sat" if unit == "sat"
+15 -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(
@@ -129,6 +139,10 @@ async def test_capped_payout_recovers_all_change_with_real_cashu_sdk(
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=liability * 1000), AsyncMock(return_value=liability * 1000),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=liability * 1000),
),
patch( patch(
"routstr.payment.lnurl.get_lnurl_data", "routstr.payment.lnurl.get_lnurl_data",
AsyncMock( AsyncMock(
+195
View File
@@ -0,0 +1,195 @@
"""Owner payout keeps each wallet's declared liability and the global total.
Regression for multi-mint payout starvation: subtracting the *total* user
liability from every wallet hid the owner surplus on any mint holding less
than the whole liability, so only the largest wallet could ever pay out.
"""
from collections.abc import AsyncIterator, Iterator
from contextlib import asynccontextmanager, contextmanager
from unittest.mock import AsyncMock, Mock, patch
import pytest
from routstr.core.settings import settings
from routstr.wallet import _owner_balance_for_mint_and_unit, _payout_mint_and_unit
MINT_A = "https://a.test"
MINT_B = "https://b.test"
@asynccontextmanager
async def _session() -> AsyncIterator[Mock]:
yield Mock()
def _keyset_id(mint_url: str, unit: str) -> str:
return f"{mint_url}|{unit}"
@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
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 (
_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:
with (
_liabilities({MINT_A: 216, MINT_B: 34}, total_sats=250),
_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
@pytest.mark.asyncio
async def test_owner_balance_never_exceeds_global_surplus() -> None:
"""Liability nobody declared against a mint is still covered in aggregate."""
with (
_liabilities({}, total_sats=250),
_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_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.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
@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."""
with (
_liabilities({}, total_sats=600),
patch.object(settings, "cashu_mints", [MINT_A, MINT_A, MINT_B]),
_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_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),
_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("load") is False for c in get_wallet.await_args_list)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"mint_liability,total_liability,expected",
[(1_500, 2_200, 1_800), (2_700, 1_500, 1_300)],
)
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."""
with (
_env(),
_wallet_db({}),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=mint_liability),
),
patch(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=total_liability),
),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "msat", 4_000) == expected
@pytest.mark.asyncio
async def test_payout_sends_the_smaller_wallets_surplus() -> None:
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),
_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),
),
patch("routstr.wallet.asyncio.sleep", AsyncMock()),
patch("routstr.wallet.raw_send_to_lnurl", send),
):
await _payout_mint_and_unit(MINT_B, "sat")
assert send.await_args is not None
assert send.await_args.kwargs["amount"] == 236
+7 -2
View File
@@ -1,6 +1,6 @@
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, call, patch
import pytest import pytest
@@ -39,11 +39,16 @@ async def test_payout_limits_and_proof_refresh(
patch( patch(
"routstr.wallet.db.total_user_liability", AsyncMock(return_value=liability) "routstr.wallet.db.total_user_liability", AsyncMock(return_value=liability)
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=liability),
),
patch("routstr.wallet.asyncio.sleep", sleep), patch("routstr.wallet.asyncio.sleep", sleep),
patch("routstr.wallet.raw_send_to_lnurl", send), patch("routstr.wallet.raw_send_to_lnurl", send),
): ):
await _payout_mint_and_unit("https://mint.test", unit) await _payout_mint_and_unit("https://mint.test", unit)
get_wallet.assert_awaited_once_with( # Later awaits belong to the other-wallet scan, which forces a reload too.
assert get_wallet.await_args_list[0] == call(
"https://mint.test", unit, force_reload_proofs=True "https://mint.test", unit, force_reload_proofs=True
) )
if expected is None: if expected is None:
+46 -7
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(
@@ -107,6 +117,10 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.db.record_lightning_payout", record_payout), patch("routstr.wallet.db.record_lightning_payout", record_payout),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout), patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
@@ -181,6 +195,10 @@ async def test_periodic_payout_releases_session_before_slow_mint_send() -> None:
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)), patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)),
): ):
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
@@ -194,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()
@@ -229,6 +245,10 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
): ):
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
@@ -236,10 +256,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.
good_calls = [ good_reloads = [
c for c in get_wallet.await_args_list if c.args[0] == "http://good:3338" c
for c in get_wallet.await_args_list
if c.args[0] == "http://good:3338" and c.kwargs.get("force_reload_proofs")
] ]
assert len(good_calls) == 2 # sat + msat # 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 assert raw_send.await_count == 2 # good mint paid for both units
@@ -331,6 +354,10 @@ async def test_periodic_payout_caps_amount_at_max_payout_sat() -> None:
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=0), AsyncMock(return_value=0),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch("routstr.wallet.raw_send_to_lnurl", raw_send), patch("routstr.wallet.raw_send_to_lnurl", raw_send),
): ):
with pytest.raises(_LoopBreak): with pytest.raises(_LoopBreak):
@@ -379,6 +406,10 @@ async def test_payout_history_records_the_capped_amount() -> None:
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
), ),
patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)), patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch( patch(
"routstr.wallet.db.list_unsettled_lightning_payouts", "routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=[]), AsyncMock(return_value=[]),
@@ -446,6 +477,10 @@ async def test_payout_history_marks_failed_only_on_proven_non_payment(
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
), ),
patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)), patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch( patch(
"routstr.wallet.db.list_unsettled_lightning_payouts", "routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=[]), AsyncMock(return_value=[]),
@@ -503,6 +538,10 @@ async def test_payout_history_write_failure_does_not_block_payout() -> None:
AsyncMock(side_effect=lambda proofs, wallet: proofs), AsyncMock(side_effect=lambda proofs, wallet: proofs),
), ),
patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)), patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=0)),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=0),
),
patch( patch(
"routstr.wallet.db.list_unsettled_lightning_payouts", "routstr.wallet.db.list_unsettled_lightning_payouts",
AsyncMock(return_value=[]), AsyncMock(return_value=[]),
@@ -0,0 +1,131 @@
"""Real-DB coverage for the per-mint liability query that bounds owner payout."""
from typing import AsyncGenerator
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, Refund, user_liability_for_mint_and_unit
MINT = "http://m1"
def _make_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def session() -> "AsyncGenerator[AsyncSession, None]":
engine = _make_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
db_session = AsyncSession(engine, expire_on_commit=False)
try:
yield db_session
finally:
await db_session.close()
await engine.dispose()
async def _add_key(
session: AsyncSession,
hashed_key: str,
balance: int,
mint_url: str | None = MINT,
currency: str | None = "sat",
) -> None:
session.add(
ApiKey(
hashed_key=hashed_key,
balance=balance,
refund_mint_url=mint_url,
refund_currency=currency,
)
)
await session.commit()
async def _add_refund(
session: AsyncSession,
hashed_key: str,
amount_msats: int,
status: str,
mint_url: str = MINT,
unit: str = "sat",
) -> None:
session.add(
Refund(
api_key_hashed_key=hashed_key,
method="lightning",
amount_msats=amount_msats,
unit=unit,
mint_url=mint_url,
status=status,
)
)
await session.commit()
@pytest.mark.asyncio
async def test_sums_key_balances_for_the_mint_and_unit(session: AsyncSession) -> None:
await _add_key(session, "a", 1000)
await _add_key(session, "b", 500)
assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1500
@pytest.mark.asyncio
async def test_adds_unresolved_refunds_to_key_balances(session: AsyncSession) -> None:
# 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)
await _add_refund(session, "b", 40, "ambiguous")
await _add_key(session, "c", 0)
await _add_refund(session, "c", 7, "stuck")
assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1347
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["paid", "failed"])
async def test_excludes_resolved_refunds(session: AsyncSession, status: str) -> None:
await _add_key(session, "a", 0)
await _add_refund(session, "a", 900, status)
assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 0
@pytest.mark.asyncio
async def test_excludes_other_mints_and_units(session: AsyncSession) -> None:
await _add_key(session, "a", 1000)
await _add_key(session, "other-mint", 111, mint_url="http://m2")
await _add_key(session, "other-unit", 222, currency="msat")
await _add_refund(session, "a", 300, "pending")
await _add_refund(session, "other-mint", 444, "pending", mint_url="http://m2")
await _add_refund(session, "other-unit", 555, "pending", unit="msat")
assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1300
@pytest.mark.asyncio
async def test_excludes_keys_without_a_refund_mint(session: AsyncSession) -> None:
await _add_key(session, "a", 1000)
await _add_key(session, "unattributed", 4242, mint_url=None, currency=None)
assert await user_liability_for_mint_and_unit(session, MINT, "sat") == 1000
@pytest.mark.asyncio
async def test_unknown_mint_has_no_liability(session: AsyncSession) -> None:
await _add_key(session, "a", 1000)
await _add_refund(session, "a", 300, "pending")
assert await user_liability_for_mint_and_unit(session, "http://missing", "sat") == 0
+64 -1
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
@@ -1254,9 +1307,14 @@ async def test_prepare_bolt11_payment_does_not_spend_user_liabilities() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() -> None: async def test_prepare_bolt11_payment_floors_fractional_owner_surplus() -> None:
"""A sub-sat surplus is not enough to fund a 1 sat invoice."""
from routstr.core.settings import settings from routstr.core.settings import settings
@asynccontextmanager
async def session() -> AsyncIterator[MagicMock]:
yield MagicMock()
wallet = MagicMock() wallet = MagicMock()
wallet.proofs = [MagicMock(amount=100)] wallet.proofs = [MagicMock(amount=100)]
wallet.melt_quote = AsyncMock( wallet.melt_quote = AsyncMock(
@@ -1281,10 +1339,15 @@ async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() -
"routstr.wallet.slow_filter_spend_proofs", "routstr.wallet.slow_filter_spend_proofs",
side_effect=lambda proofs, wallet: proofs, side_effect=lambda proofs, wallet: proofs,
), ),
patch("routstr.wallet.db.create_session", session),
patch( patch(
"routstr.wallet.db.total_user_liability", "routstr.wallet.db.total_user_liability",
AsyncMock(return_value=99_999), AsyncMock(return_value=99_999),
), ),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=99_999),
),
pytest.raises(ValueError, match="user liabilities"), pytest.raises(ValueError, match="user liabilities"),
): ):
await prepare_bolt11_payment("lnbc-invoice") await prepare_bolt11_payment("lnbc-invoice")