mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: read other wallets' proofs fresh when bounding owner payout
This commit is contained in:
+1
-2
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user