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)
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(
db_session: AsyncSession, mint_url: str, unit: str
) -> int:
+98 -35
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
@@ -121,7 +122,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
@@ -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):
@@ -701,18 +705,64 @@ class Bolt11PaymentPlan:
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(
mint_url: str, unit: str, proofs_balance: 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:
# Refund mint is a preference, not funding provenance. Mirror payout's
# conservative rule and protect the full liability at every mint.
user_liability = await db.total_user_liability(session)
# API-key balances are stored in msats. Cashu ``sat`` proofs are not.
if unit == "sat":
user_liability = _msats_to_sats_ceil(user_liability)
return max(0, proofs_balance - user_liability)
mint_liability = await db.user_liability_for_mint_and_unit(
session, mint_url, unit
)
total_liability = await db.total_user_liability(session)
proofs_msats = _to_msats(proofs_balance, unit)
surplus_msats = min(
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:
@@ -787,9 +837,7 @@ async def _prepare_bolt11_payment(invoice: str) -> Bolt11PaymentPlan:
)
if owner_balance < required:
continue
owner_balance_msats = (
owner_balance * 1000 if unit == "sat" else owner_balance
)
owner_balance_msats = _to_msats(owner_balance, unit)
candidates.append(
(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.
_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] = {}
@@ -1189,16 +1240,34 @@ async def get_wallet(
or last_mint_load is None
or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
):
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,
)
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)
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()
if load_proofs:
@@ -1650,15 +1719,13 @@ async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
)
return
# Fetch liability after the proofs snapshot and settle delay while the
# wallet operation guard excludes concurrent proof mutation and crediting.
# Read liabilities and the other wallets' proofs after this wallet's proofs
# snapshot and settle delay, while the wallet operation guard excludes
# concurrent proof mutation and crediting.
try:
async with db.create_session() as session:
# ApiKey stores a refund preference, not funding provenance. Until
# 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)
available_balance = await _owner_balance_for_mint_and_unit(
mint_url, unit, sum(proof.amount for proof in proofs)
)
except Exception as e:
logger.error(
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
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 = (
settings.max_payout_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 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(
@@ -129,6 +139,10 @@ async def test_capped_payout_recovers_all_change_with_real_cashu_sdk(
"routstr.wallet.db.total_user_liability",
AsyncMock(return_value=liability * 1000),
),
patch(
"routstr.wallet.db.user_liability_for_mint_and_unit",
AsyncMock(return_value=liability * 1000),
),
patch(
"routstr.payment.lnurl.get_lnurl_data",
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 contextlib import asynccontextmanager
from unittest.mock import AsyncMock, Mock, patch
from unittest.mock import AsyncMock, Mock, call, patch
import pytest
@@ -39,11 +39,16 @@ async def test_payout_limits_and_proof_refresh(
patch(
"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.raw_send_to_lnurl", send),
):
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
)
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).
"""
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(
@@ -107,6 +117,10 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
"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("routstr.wallet.db.record_lightning_payout", record_payout),
patch("routstr.wallet.db.settle_lightning_payout", settle_payout),
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",
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)),
):
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."""
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()
@@ -229,6 +245,10 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
"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("routstr.wallet.raw_send_to_lnurl", raw_send),
):
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
# 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"
good_reloads = [
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
@@ -331,6 +354,10 @@ async def test_periodic_payout_caps_amount_at_max_payout_sat() -> None:
"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("routstr.wallet.raw_send_to_lnurl", raw_send),
):
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),
),
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(
"routstr.wallet.db.list_unsettled_lightning_payouts",
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),
),
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(
"routstr.wallet.db.list_unsettled_lightning_payouts",
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),
),
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(
"routstr.wallet.db.list_unsettled_lightning_payouts",
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._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
@@ -1254,9 +1307,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 +1339,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")