fix: prevent duplicate foreign swap credits

This commit is contained in:
9qeklajc
2026-10-04 22:06:40 +02:00
parent eed132b422
commit 3301a8fb8b
2 changed files with 113 additions and 23 deletions
+52 -2
View File
@@ -239,6 +239,30 @@ async def _update(swap: CashuSwap, **values: Any) -> None:
setattr(swap, name, value) setattr(swap, name, value)
async def _transition_status(
swap: CashuSwap, expected: str, status: str, **values: Any
) -> bool:
values.update(status=status, updated_at=int(time.time()))
async with db.create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(CashuSwap)
.where(col(CashuSwap.id) == swap.id)
.where(col(CashuSwap.status) == expected)
.values(**values)
)
await session.commit()
transitioned = (getattr(result, "rowcount", 0) or 0) == 1
if transitioned:
for name, value in values.items():
setattr(swap, name, value)
return transitioned
async def _load_swap(swap_id: str) -> CashuSwap | None:
async with db.create_session() as session:
return await session.get(CashuSwap, swap_id)
async def _prior_swap_for_token(token_hash: str) -> CashuSwap | None: async def _prior_swap_for_token(token_hash: str) -> CashuSwap | None:
async with db.create_session() as session: async with db.create_session() as session:
result = await session.exec( result = await session.exec(
@@ -549,6 +573,11 @@ async def _finish_swap_in(
session: AsyncSession | None = None, session: AsyncSession | None = None,
) -> int: ) -> int:
"""Mint on the trusted destination and credit the key, under the guard.""" """Mint on the trusted destination and credit the key, under the guard."""
started_melted = swap.status == "melted"
def credited_msats() -> int:
return swap.destination_amount * (1000 if swap.destination_unit == "sat" else 1)
async with wallet_operation_guard(): async with wallet_operation_guard():
if swap.status == "melted": if swap.status == "melted":
dest_wallet = await get_wallet(swap.destination_mint, swap.destination_unit) dest_wallet = await get_wallet(swap.destination_mint, swap.destination_unit)
@@ -574,11 +603,23 @@ async def _finish_swap_in(
raise SwapPendingError( raise SwapPendingError(
"Destination mint failed; retrying later" "Destination mint failed; retrying later"
) from error ) from error
await _update(swap, status="minted", error=None) transitioned = await _transition_status(
swap, "melted", "minted", error=None
)
if not transitioned:
stored = await _load_swap(swap.id)
if stored is not None and stored.status == "credited":
return credited_msats()
if stored is None or stored.status != "minted":
status = stored.status if stored is not None else "missing"
raise SwapPendingError(f"Swap is {status}")
swap.status = "minted"
swap.error = stored.error
if swap.status != "minted": if swap.status != "minted":
raise SwapPendingError(f"Swap is {swap.status}") raise SwapPendingError(f"Swap is {swap.status}")
try:
if session is None or key is None: if session is None or key is None:
async with db.create_session() as own_session: async with db.create_session() as own_session:
own_key = await own_session.get(ApiKey, swap.api_key_hashed_key) own_key = await own_session.get(ApiKey, swap.api_key_hashed_key)
@@ -586,7 +627,10 @@ async def _finish_swap_in(
await _update(swap, status="failed", error="api key missing") await _update(swap, status="failed", error="api key missing")
logger.critical( logger.critical(
"Swapped funds have no API key to credit", "Swapped funds have no API key to credit",
extra={"swap_id": swap.id, "amount": swap.destination_amount}, extra={
"swap_id": swap.id,
"amount": swap.destination_amount,
},
) )
raise TokenConsumedError("API key vanished before swap credit") raise TokenConsumedError("API key vanished before swap credit")
credited = await _apply_credit_locked( credited = await _apply_credit_locked(
@@ -610,6 +654,12 @@ async def _finish_swap_in(
refund_mint_url=swap.source_mint, refund_mint_url=swap.source_mint,
swap_id=swap.id, swap_id=swap.id,
) )
except TokenConsumedError:
if started_melted:
stored = await _load_swap(swap.id)
if stored is not None and stored.status == "credited":
return credited_msats()
raise
swap.status = "credited" swap.status = "credited"
swap.error = None swap.error = None
logger.info( logger.info(
+40
View File
@@ -539,6 +539,46 @@ async def test_minted_swap_credit_is_atomic_and_cannot_repeat(
assert stored.status == "credited" assert stored.status == "credited"
@pytest.mark.asyncio
async def test_stale_melted_worker_cannot_reset_and_credit_swap_twice(
engine: AsyncEngine, session: AsyncSession
) -> None:
key = await _make_key(session)
swap = CashuSwap(
direction="in",
status="melted",
api_key_hashed_key=key.hashed_key,
token="cashuAstalemelted",
token_hash="stale-melted",
source_mint=FOREIGN,
source_unit="sat",
source_amount=1000,
destination_mint=PRIMARY,
destination_unit="sat",
destination_amount=998,
mint_quote_id="mint-998",
melt_quote_id="melt-998",
)
await fms._save(swap)
async with AsyncSession(session.bind, expire_on_commit=False) as stale_session:
stale_swap = await stale_session.get(CashuSwap, swap.id)
assert stale_swap is not None
foreign = _ForeignWallet()
primary = _PrimaryWallet()
async with _swap_env(foreign, primary, _token()):
first = await fms._finish_swap_in(swap, key=key, session=session)
second = await fms._finish_swap_in(stale_swap, key=key, session=session)
assert first == second == 998_000
await session.refresh(key)
assert key.balance == 998_000
rows = list((await session.exec(select(CashuTransaction))).all())
assert len(rows) == 1
(stored,) = await _swap_rows(session)
assert stored.status == "credited"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mint_recovery_reuses_proofs_tagged_with_quote() -> None: async def test_mint_recovery_reuses_proofs_tagged_with_quote() -> None:
proof = SimpleNamespace(amount=998, reserved=False, mint_id="mint-998") proof = SimpleNamespace(amount=998, reserved=False, mint_id="mint-998")