fix: read other wallets' proofs fresh when bounding owner payout

This commit is contained in:
9qeklajc
2026-09-23 19:42:49 +02:00
parent bb40e869e6
commit b663b1d83a
7 changed files with 214 additions and 38 deletions
+1 -2
View File
@@ -1103,8 +1103,7 @@ async def user_liability_for_mint_and_unit(
) -> 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``.
Single statement, for the same atomicity reason as ``total_user_liability``.
"""
key_balances = (
select(func.coalesce(func.sum(ApiKey.balance), 0))
+14 -18
View File
@@ -121,7 +121,7 @@ def _msats_to_sats_ceil(amount: int) -> int:
def _mints_to_inspect() -> list[str]:
"""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:
mint_urls.append(settings.primary_mint)
return mint_urls
@@ -708,8 +708,10 @@ 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.
Reads local proof snapshots only. A wallet that cannot be loaded counts
as empty, which can only shrink the owner surplus.
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.
"""
total = 0
for other_mint in _mints_to_inspect():
@@ -717,16 +719,11 @@ async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int:
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),
},
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
@@ -738,13 +735,12 @@ async def _other_wallets_unreserved_msats(mint_url: str, unit: str) -> int:
async def _owner_balance_for_mint_and_unit(
mint_url: str, unit: str, proofs_balance: int
) -> int:
"""Return owner funds in one wallet, in that wallet's unit, never negative.
"""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. 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.
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:
@@ -757,7 +753,7 @@ async def _owner_balance_for_mint_and_unit(
proofs_msats - mint_liability,
proofs_msats + others_msats - total_liability,
)
# Cashu ``sat`` proofs are whole sats; round the surplus down, never up.
# Cashu ``sat`` proofs are whole sats.
surplus = _msats_to_sats(surplus_msats) if unit == "sat" else surplus_msats
return max(0, surplus)
+42 -6
View File
@@ -102,24 +102,60 @@ async def test_unloadable_wallet_counts_as_empty() -> None:
@pytest.mark.asyncio
async def test_msat_wallet_surplus_is_not_rounded() -> None:
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),
):
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})
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),
):
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)
@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."""
get_wallet, get_proofs = _wallets({MINT_A: 0, MINT_B: 0})
with (
_liabilities({MINT_B: 0}, total_sats=0),
_liabilities({}, 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),
AsyncMock(return_value=mint_liability),
),
patch(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=total_liability),
),
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
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:
"""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 (
+4 -4
View File
@@ -47,10 +47,10 @@ async def test_payout_limits_and_proof_refresh(
patch("routstr.wallet.raw_send_to_lnurl", send),
):
await _payout_mint_and_unit("https://mint.test", unit)
reloads = [
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)]
# 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
)
if expected is None:
send.assert_not_awaited()
else:
+6 -7
View File
@@ -14,7 +14,7 @@ from collections.abc import Callable, Coroutine
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Any
from unittest.mock import ANY, AsyncMock, MagicMock, patch
from unittest.mock import ANY, AsyncMock, MagicMock, call, patch
import pytest
@@ -248,12 +248,11 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
# 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.
good_calls = [
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
for unit in ("sat", "msat"):
assert (
call("http://good:3338", unit, force_reload_proofs=True)
in get_wallet.await_args_list
)
assert raw_send.await_count == 2 # good mint paid for both units
@@ -0,0 +1,136 @@
"""Real-DB coverage for db.user_liability_for_mint_and_unit.
Verifies the per-mint liability query that bounds owner payout: it sums key
balances and unresolved refund claims for one (mint_url, unit), excludes
resolved claims and other mints/units, and drops keys with no refund mint.
"""
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:
# One open claim per key, so each unresolved status needs its own key.
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
+11 -1
View File
@@ -1254,9 +1254,14 @@ async def test_prepare_bolt11_payment_does_not_spend_user_liabilities() -> None:
@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
@asynccontextmanager
async def session() -> AsyncIterator[MagicMock]:
yield MagicMock()
wallet = MagicMock()
wallet.proofs = [MagicMock(amount=100)]
wallet.melt_quote = AsyncMock(
@@ -1281,10 +1286,15 @@ async def test_prepare_bolt11_payment_rounds_user_liability_up_to_whole_sats() -
"routstr.wallet.slow_filter_spend_proofs",
side_effect=lambda proofs, wallet: proofs,
),
patch("routstr.wallet.db.create_session", session),
patch(
"routstr.wallet.db.total_user_liability",
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"),
):
await prepare_bolt11_payment("lnbc-invoice")