fix: bound owner payout by per-mint liability instead of total

This commit is contained in:
9qeklajc
2026-09-23 17:32:31 +02:00
parent 50d5601547
commit bb40e869e6
6 changed files with 272 additions and 24 deletions
+29
View File
@@ -1098,6 +1098,35 @@ 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.
Key balances and unresolved refund claims are summed in one statement for
the same 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:
+59 -19
View File
@@ -701,18 +701,65 @@ 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.
Reads local proof snapshots only. A wallet that cannot be loaded counts
as empty, which can only shrink the owner surplus.
"""
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)
except Exception as e:
logger.debug(
"Wallet excluded from owner surplus",
extra={
"mint_url": other_mint,
"unit": other_unit,
"error": str(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
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, never negative.
A key's refund mint is a preference, not funding provenance: a key topped
up from a second mint keeps its original refund mint. So two bounds apply.
The wallet keeps the liability declared against it, so refunds drawn from
it stay serviceable. All wallets together keep the total liability, so
misattributed customer funds are never 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; round the surplus down, never up.
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:
@@ -1650,15 +1697,12 @@ 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 after the proofs snapshot and settle delay while the
# wallet operation guard excludes concurrent proof mutation and crediting. # 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 +1711,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"
+4
View File
@@ -129,6 +129,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(
+140
View File
@@ -0,0 +1,140 @@
"""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 ExitStack, 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 _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."""
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 _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 ExitStack() as stack:
for target in (
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),
),
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
@pytest.mark.asyncio
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})
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),
):
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."""
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),
):
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:
"""A wallet that cannot be read shrinks the surplus rather than inflating it."""
get_wallet, get_proofs = _wallets(
{MINT_A: 400, MINT_B: 270}, unreachable=frozenset({MINT_A})
)
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),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "sat", 270) == 20
@pytest.mark.asyncio
async def test_msat_wallet_surplus_is_not_rounded() -> None:
get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0})
with (
_liabilities({MINT_B: 0}, total_sats=0),
patch("routstr.wallet.get_wallet", get_wallet),
patch("routstr.wallet.get_proofs_per_mint_and_unit", get_proofs),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=1_500),
),
patch("routstr.wallet.db.total_user_liability", AsyncMock(return_value=1_500)),
):
assert await _owner_balance_for_mint_and_unit(MINT_B, "msat", 4_000) == 2_500
@pytest.mark.asyncio
async def test_payout_sends_the_smaller_wallets_surplus() -> None:
"""End to end: the wallet below the total liability still pays out."""
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),
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
+9 -4
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,13 +39,18 @@ 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( reloads = [
"https://mint.test", unit, force_reload_proofs=True c for c in get_wallet.await_args_list if c.kwargs.get("force_reload_proofs")
) ]
assert reloads == [call("https://mint.test", unit, force_reload_proofs=True)]
if expected is None: if expected is None:
send.assert_not_awaited() send.assert_not_awaited()
else: else:
+31 -1
View File
@@ -107,6 +107,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 +185,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):
@@ -229,6 +237,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):
@@ -237,7 +249,9 @@ 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_calls = [
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 assert len(good_calls) == 2 # sat + msat
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 +345,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 +397,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 +468,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 +529,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=[]),