diff --git a/routstr/foreign_mint_swap.py b/routstr/foreign_mint_swap.py index 57541078..6f832536 100644 --- a/routstr/foreign_mint_swap.py +++ b/routstr/foreign_mint_swap.py @@ -239,6 +239,30 @@ async def _update(swap: CashuSwap, **values: Any) -> None: 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 with db.create_session() as session: result = await session.exec( @@ -549,6 +573,11 @@ async def _finish_swap_in( session: AsyncSession | None = None, ) -> int: """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(): if swap.status == "melted": dest_wallet = await get_wallet(swap.destination_mint, swap.destination_unit) @@ -574,24 +603,50 @@ async def _finish_swap_in( raise SwapPendingError( "Destination mint failed; retrying later" ) 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": raise SwapPendingError(f"Swap is {swap.status}") - if session is None or key is None: - async with db.create_session() as own_session: - own_key = await own_session.get(ApiKey, swap.api_key_hashed_key) - if own_key is None: - await _update(swap, status="failed", error="api key missing") - logger.critical( - "Swapped funds have no API key to credit", - extra={"swap_id": swap.id, "amount": swap.destination_amount}, + try: + if session is None or key is None: + async with db.create_session() as own_session: + own_key = await own_session.get(ApiKey, swap.api_key_hashed_key) + if own_key is None: + await _update(swap, status="failed", error="api key missing") + logger.critical( + "Swapped funds have no API key to credit", + extra={ + "swap_id": swap.id, + "amount": swap.destination_amount, + }, + ) + raise TokenConsumedError("API key vanished before swap credit") + credited = await _apply_credit_locked( + own_key, + own_session, + amount=swap.destination_amount, + unit=swap.destination_unit, + mint_url=swap.destination_mint, + token=str(swap.token), + refund_mint_url=swap.source_mint, + swap_id=swap.id, ) - raise TokenConsumedError("API key vanished before swap credit") + else: credited = await _apply_credit_locked( - own_key, - own_session, + key, + session, amount=swap.destination_amount, unit=swap.destination_unit, mint_url=swap.destination_mint, @@ -599,17 +654,12 @@ async def _finish_swap_in( refund_mint_url=swap.source_mint, swap_id=swap.id, ) - else: - credited = await _apply_credit_locked( - key, - session, - amount=swap.destination_amount, - unit=swap.destination_unit, - mint_url=swap.destination_mint, - token=str(swap.token), - refund_mint_url=swap.source_mint, - 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.error = None logger.info( diff --git a/tests/unit/test_foreign_mint_swap.py b/tests/unit/test_foreign_mint_swap.py index f995c7ad..84852737 100644 --- a/tests/unit/test_foreign_mint_swap.py +++ b/tests/unit/test_foreign_mint_swap.py @@ -539,6 +539,46 @@ async def test_minted_swap_credit_is_atomic_and_cannot_repeat( 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 async def test_mint_recovery_reuses_proofs_tagged_with_quote() -> None: proof = SimpleNamespace(amount=998, reserved=False, mint_id="mint-998")