diff --git a/migrations/versions/f3a1c7b9e2d4_add_refunds_table.py b/migrations/versions/f3a1c7b9e2d4_add_refunds_table.py new file mode 100644 index 00000000..e0b3c823 --- /dev/null +++ b/migrations/versions/f3a1c7b9e2d4_add_refunds_table.py @@ -0,0 +1,58 @@ +"""add refunds table + +Revision ID: f3a1c7b9e2d4 +Revises: e5a6b7c8d9f0 +Create Date: 2026-09-07 + +""" + +import sqlalchemy as sa +import sqlmodel +from alembic import op + +revision = "f3a1c7b9e2d4" +down_revision = "e5a6b7c8d9f0" +branch_labels = None +depends_on = None + +OPEN_STATUSES = "status IN ('pending', 'ambiguous')" + + +def upgrade() -> None: + op.create_table( + "refunds", + sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column( + "api_key_hashed_key", sqlmodel.sql.sqltypes.AutoString(), nullable=False + ), + sa.Column("method", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("destination", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("amount_msats", sa.Integer(), nullable=False), + sa.Column("unit", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("status", sqlmodel.sql.sqltypes.AutoString(), nullable=False), + sa.Column("quote_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("token", sqlmodel.sql.sqltypes.AutoString(), nullable=True), + sa.Column("claimed_at", sa.Integer(), nullable=True), + sa.Column("created_at", sa.Integer(), nullable=False), + sa.Column("updated_at", sa.Integer(), nullable=False), + sa.ForeignKeyConstraint(["api_key_hashed_key"], ["api_keys.hashed_key"]), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("ix_refunds_api_key_hashed_key", "refunds", ["api_key_hashed_key"]) + op.create_index("ix_refunds_status", "refunds", ["status"]) + op.create_index( + "ux_refunds_open_per_key", + "refunds", + ["api_key_hashed_key"], + unique=True, + sqlite_where=sa.text(OPEN_STATUSES), + postgresql_where=sa.text(OPEN_STATUSES), + ) + + +def downgrade() -> None: + op.drop_index("ux_refunds_open_per_key", table_name="refunds") + op.drop_index("ix_refunds_status", table_name="refunds") + op.drop_index("ix_refunds_api_key_hashed_key", table_name="refunds") + op.drop_table("refunds") diff --git a/routstr/balance.py b/routstr/balance.py index ecd770b5..7c37331a 100644 --- a/routstr/balance.py +++ b/routstr/balance.py @@ -1,13 +1,12 @@ -import asyncio import hashlib -from time import monotonic from typing import Annotated, NoReturn from fastapi import APIRouter, Depends, Header, HTTPException from fastapi.responses import JSONResponse from pydantic import BaseModel -from sqlmodel import col, select, update +from sqlmodel import col, select +from . import refund from .auth import ( redemption_error_to_http_exception, validate_bearer_key, @@ -19,20 +18,13 @@ from .core.db import ( get_session, release_stale_reservations, ) -from .core.db import ( - store_cashu_transaction_with_retry as store_cashu_transaction, -) from .core.logging import get_logger from .core.settings import settings from .lightning import lightning_router -from .payment.lnurl import MeltOutcomeAmbiguousError from .wallet import ( classify_redemption_error, credit_balance, - is_mint_connection_error, recieve_token, - send_to_lnurl, - send_token, token_mint_url, ) @@ -236,35 +228,6 @@ async def topup_wallet_endpoint( return {"msats": amount_msats} -_REFUND_CACHE_TTL_SECONDS: int = settings.refund_cache_ttl_seconds -_refund_cache_lock: asyncio.Lock = asyncio.Lock() -_refund_cache: dict[str, tuple[float, dict[str, str]]] = {} - - -def _cache_key_for_authorization(authorization: str) -> str: - return hashlib.sha256(authorization.strip().encode()).hexdigest() - - -async def _refund_cache_get(authorization: str) -> dict[str, str] | None: - key = _cache_key_for_authorization(authorization) - async with _refund_cache_lock: - item = _refund_cache.get(key) - if item is None: - return None - expires_at, value = item - if expires_at <= monotonic(): - del _refund_cache[key] - return None - return value - - -async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None: - key = _cache_key_for_authorization(authorization) - expiry = monotonic() + _REFUND_CACHE_TTL_SECONDS - async with _refund_cache_lock: - _refund_cache[key] = (expiry, value) - - async def _lookup_key_no_create( bearer_value: str, session: AsyncSession ) -> ApiKey | None: @@ -307,36 +270,13 @@ async def _get_persisted_api_key_refund( return persisted -async def _restore_balance( - session: AsyncSession, - hashed_key: str, - balance: int, - reserved_balance: int, - mint_url: str, -) -> None: - """Restore balance after a failed refund mint attempt.""" - restore_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == hashed_key) - .values( - balance=col(ApiKey.balance) + balance, - reserved_balance=col(ApiKey.reserved_balance) + reserved_balance, - ) - ) - await session.exec(restore_stmt) # type: ignore[call-overload] - await session.commit() - logger.info( - "refund_wallet_endpoint: balance restored after mint failure", - extra={ - "key_hash": hashed_key[:8], - "restored_balance": balance, - "mint_url": mint_url, - }, - ) +class RefundRequest(BaseModel): + lightning_address: str | None = None @router.post("/refund", response_model=None) async def refund_wallet_endpoint( + refund_request: RefundRequest | None = None, authorization: Annotated[str | None, Header()] = None, x_cashu: Annotated[str | None, Header()] = None, session: AsyncSession = Depends(get_session), @@ -412,9 +352,14 @@ async def refund_wallet_endpoint( }, ) + requested = refund_request.lightning_address if refund_request else None + if requested: + await refund.validate_lightning_destination(requested) + destination = requested or key.refund_address + if key.total_balance <= 0: - if cached := await _refund_cache_get(bearer_value): - return cached + if paid := await refund.latest_terminal(session, key): + return refund.describe(paid) if persisted := await _get_persisted_api_key_refund(key, session): return persisted @@ -441,162 +386,21 @@ async def refund_wallet_endpoint( ) remaining_balance_msats: int = key.total_balance - - if key.refund_currency == "sat": - remaining_balance = remaining_balance_msats // 1000 - else: - remaining_balance = remaining_balance_msats + unit = refund.refund_unit(key) + remaining_balance = refund.amount_in_unit(remaining_balance_msats, unit) if remaining_balance_msats > 0 and remaining_balance <= 0: raise HTTPException(status_code=400, detail="Balance too small to refund") elif remaining_balance <= 0: raise HTTPException(status_code=400, detail="No balance to refund") - # Capture values before debit — the session may refresh key after commit - pre_debit_balance = key.balance - pre_debit_reserved = key.reserved_balance - - # --- DEBIT FIRST: atomically zero the balance before minting tokens --- - # This prevents the race where a concurrent topup/spend happens between - # reading the balance and minting the refund token (double-spend). - debit_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) == pre_debit_balance) - .where(col(ApiKey.reserved_balance) == pre_debit_reserved) - .values(balance=0, reserved_balance=0, reserved_at=None) + claim = await refund.open_claim( + session, + key, + method="lightning" if destination else "cashu", + destination=destination, ) - debit_result = await session.exec(debit_stmt) # type: ignore[call-overload] - await session.commit() - - if debit_result.rowcount == 0: - # Balance changed between read and debit — another request is active - raise HTTPException( - status_code=409, - detail="Balance changed concurrently. Please retry the refund.", - ) - - # The balance is locked at zero, so it is safe to create the refund token. - effective_refund_mint = ( - key.refund_mint_url - if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints - else settings.primary_mint - ) - try: - refund_currency = key.refund_currency or "sat" - if key.refund_address: - await send_to_lnurl( - remaining_balance, - key.refund_currency or "sat", - effective_refund_mint, - key.refund_address, - ) - result = {"recipient": key.refund_address} - else: - token = await send_token( - remaining_balance, refund_currency, effective_refund_mint - ) - effective_refund_mint = token_mint_url(token, effective_refund_mint) - result = {"token": token} - - if key.refund_currency == "sat": - result["sats"] = str(remaining_balance_msats // 1000) - else: - result["msats"] = str(remaining_balance_msats) - - if "token" in result: - logger.info( - "refund_wallet_endpoint: cashu token issued", - extra={ - "path": "/v1/wallet/refund", - "token_length": len(result["token"]), - "amount": remaining_balance, - "currency": key.refund_currency or "sat", - }, - ) - - except MeltOutcomeAmbiguousError as e: - # The melt was dispatched and may still settle. Restoring the balance - # here would let the same debit be paid out twice; keep the debit and - # leave the outcome to reconciliation. - logger.error( - "refund_wallet_endpoint: melt outcome ambiguous; balance withheld " - "pending reconciliation", - extra={ - "error": str(e), - "key_hash": key.hashed_key[:8], - "remaining_balance": remaining_balance, - "refund_currency": key.refund_currency, - "refund_mint_url": key.refund_mint_url, - }, - ) - raise HTTPException( - status_code=502, - detail=( - "Refund was dispatched but its outcome is unconfirmed; the " - "balance is withheld until reconciliation completes" - ), - ) - except HTTPException: - # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) - raise - except Exception as e: - # Minting failed — restore the debited balance - await _restore_balance( - session, - key.hashed_key, - pre_debit_balance, - pre_debit_reserved, - key.refund_mint_url or "", - ) - error_msg = str(e) - logger.error( - "refund_wallet_endpoint: mint/send failed", - extra={ - "error": error_msg, - "error_type": type(e).__name__, - "key_hash": key.hashed_key[:8], - "remaining_balance": remaining_balance, - "refund_currency": key.refund_currency, - "refund_mint_url": key.refund_mint_url, - "has_refund_address": bool(key.refund_address), - }, - ) - if is_mint_connection_error(e): - raise HTTPException(status_code=503, detail="Mint service unavailable") - else: - raise HTTPException(status_code=500, detail="Refund failed") - - await _refund_cache_set(bearer_value, result) - - if "token" in result: - await store_cashu_transaction( - token=result["token"], - amount=remaining_balance, - unit=key.refund_currency or "sat", - mint_url=effective_refund_mint, - typ="out", - collected=False, - source="apikey", - api_key_hashed_key=key.hashed_key, - ) - - logger.info( - "refund_wallet_endpoint: refund successful", - extra={ - "refunded_msats": remaining_balance_msats, - "previous_reserved_balance": key.reserved_balance, - }, - ) - - return result + return await refund.execute(session, claim) @router.get("/history") @@ -640,6 +444,7 @@ async def donate(token: str, ref: str | None = None) -> str: except Exception: return "Invalid token." + @router.api_route( "/{path:path}", methods=["GET", "POST", "PUT", "DELETE"], diff --git a/routstr/core/db.py b/routstr/core/db.py index 4c4d1a62..55693d41 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -12,7 +12,7 @@ from typing import AsyncGenerator from alembic import command from alembic.config import Config from alembic.util.exc import CommandError -from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_ +from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_, text from sqlalchemy.engine import make_url from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.ext.asyncio import AsyncEngine @@ -300,9 +300,7 @@ async def release_stale_reservations( col(ApiKey.reserved_at) < cutoff ) else: - legacy_query = legacy_query.where( - col(ApiKey.hashed_key) == key_hash - ).where( + legacy_query = legacy_query.where(col(ApiKey.hashed_key) == key_hash).where( or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff) ) @@ -557,6 +555,55 @@ class CashuTransaction(SQLModel, table=True): # type: ignore ) +REFUND_OPEN_STATUSES = ("pending", "ambiguous") + +_REFUND_OPEN_PREDICATE = "status IN ('pending', 'ambiguous')" + + +class Refund(SQLModel, table=True): # type: ignore + """A durable claim on an API key's balance for a single payout. + + The partial unique index is the double-refund guarantee: a key can have at + most one open claim, so a Cashu refund cannot start while a Lightning + refund is in flight, and neither survives a crash without a record. + """ + + __tablename__ = "refunds" + __table_args__ = ( + Index( + "ux_refunds_open_per_key", + "api_key_hashed_key", + unique=True, + sqlite_where=text(_REFUND_OPEN_PREDICATE), + postgresql_where=text(_REFUND_OPEN_PREDICATE), + ), + ) + + id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex) + api_key_hashed_key: str = Field(foreign_key="api_keys.hashed_key", index=True) + method: str = Field(description="Payout method: lightning or cashu") + destination: str | None = Field( + default=None, description="Lightning address or LNURL, NULL for cashu" + ) + amount_msats: int = Field(description="Balance debited when the claim opened") + unit: str = Field(description="Mint unit the payout is denominated in") + mint_url: str = Field(description="Mint the payout is drawn from") + status: str = Field( + default="pending", + index=True, + description="pending, paid, failed, ambiguous, or stuck", + ) + quote_id: str | None = Field( + default=None, description="Melt quote id, for reconciling an ambiguous payout" + ) + token: str | None = Field(default=None, description="Issued cashu token") + claimed_at: int | None = Field( + default=None, description="Reconciler lease timestamp" + ) + created_at: int = Field(default_factory=lambda: int(time.time())) + updated_at: int = Field(default_factory=lambda: int(time.time())) + + async def store_cashu_transaction( token: str, amount: int, diff --git a/routstr/core/main.py b/routstr/core/main.py index fdb3b913..938e2144 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -28,6 +28,7 @@ from ..nostr.discovery import providers_router from ..payment.models import models_router, update_sats_pricing from ..payment.price import update_prices_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically +from ..refund import periodic_refund_reconcile from ..upstream.auto_topup import periodic_auto_topup from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing from ..upstream.litellm_routing import configure_litellm @@ -68,6 +69,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: dead_key_prune_task = None auto_topup_task = None refund_sweep_task = None + refund_reconcile_task = None routstr_fee_task = None invoice_watcher_task = None @@ -158,6 +160,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune()) auto_topup_task = asyncio.create_task(periodic_auto_topup()) refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) + refund_reconcile_task = asyncio.create_task(periodic_refund_reconcile()) routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout()) invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher()) @@ -201,6 +204,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: auto_topup_task.cancel() if refund_sweep_task is not None: refund_sweep_task.cancel() + if refund_reconcile_task is not None: + refund_reconcile_task.cancel() if routstr_fee_task is not None: routstr_fee_task.cancel() if invoice_watcher_task is not None: @@ -234,6 +239,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(auto_topup_task) if refund_sweep_task is not None: tasks_to_wait.append(refund_sweep_task) + if refund_reconcile_task is not None: + tasks_to_wait.append(refund_reconcile_task) if routstr_fee_task is not None: tasks_to_wait.append(routstr_fee_task) if invoice_watcher_task is not None: diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 971cec19..2104aeb3 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -124,7 +124,6 @@ class Settings(BaseSettings): enable_model_paths_refresh: bool = Field( default=True, env="ENABLE_MODEL_PATHS_REFRESH" ) - refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") # Uncollected refund tokens are swept after ~6 months (180 days). # Fixed for now: not configurable via env or the settings DB/admin API # (empty env list disables env binding; see FIXED_FIELDS). @@ -132,6 +131,14 @@ class Settings(BaseSettings): refund_sweep_claim_timeout_seconds: int = Field( default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS" ) + # How long an open refund claim may sit before the reconciler asks the mint + # what became of it. Doubles as the reconciler's per-row lease. + refund_claim_timeout_seconds: int = Field( + default=300, gt=0, env="REFUND_CLAIM_TIMEOUT_SECONDS" + ) + refund_reconcile_interval_seconds: int = Field( + default=60, gt=0, env="REFUND_RECONCILE_INTERVAL_SECONDS" + ) # Database connection-pool controls (advanced). Capacity defaults provide # headroom for Routstr's concurrent request and background-payment workload. diff --git a/routstr/refund.py b/routstr/refund.py new file mode 100644 index 00000000..75c2a0db --- /dev/null +++ b/routstr/refund.py @@ -0,0 +1,429 @@ +"""Durable, mutually exclusive refund claims for API key balances. + +A claim debits the balance and records the payout in one transaction, so a key +can never have two payouts in flight and no crash can leave a debited balance +without a record of why. +""" + +import asyncio +import time + +from fastapi import HTTPException +from sqlalchemy.exc import IntegrityError +from sqlmodel import col, select, update + +from .core.db import ( + REFUND_OPEN_STATUSES, + ApiKey, + AsyncSession, + Refund, + create_session, +) +from .core.db import ( + store_cashu_transaction_with_retry as store_cashu_transaction, +) +from .core.logging import get_logger +from .core.settings import settings +from .payment.lnurl import LNURLError, MeltOutcomeAmbiguousError, get_lnurl_data +from .wallet import ( + check_bolt11_payment_status, + is_mint_connection_error, + send_to_lnurl, + send_token, + token_mint_url, +) + +logger = get_logger(__name__) + +RECONCILE_BATCH_LIMIT = 100 + + +def refund_unit(key: ApiKey) -> str: + return key.refund_currency or "sat" + + +def amount_in_unit(amount_msats: int, unit: str) -> int: + return amount_msats // 1000 if unit == "sat" else amount_msats + + +def refund_mint(key: ApiKey) -> str: + """Persisted mint preferences must not outlive the trusted-mint config.""" + if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints: + return key.refund_mint_url + return settings.primary_mint + + +async def validate_lightning_destination(destination: str) -> None: + """Resolve the destination before claiming, so a bad address never debits.""" + try: + await get_lnurl_data(destination) + except LNURLError as e: + raise HTTPException( + status_code=400, detail=f"Invalid lightning destination: {e}" + ) + + +async def open_claim( + session: AsyncSession, + key: ApiKey, + *, + method: str, + destination: str | None, +) -> Refund: + """Debit the balance to zero and record the claim in one transaction. + + The claim starts leased (``claimed_at``) to the request that opened it, so + the reconciler leaves it alone until ``refund_claim_timeout_seconds`` have + passed; a payout still in flight is never released underneath itself. + """ + unit = refund_unit(key) + refund = Refund( + api_key_hashed_key=key.hashed_key, + method=method, + destination=destination, + amount_msats=key.total_balance, + unit=unit, + mint_url=refund_mint(key), + claimed_at=int(time.time()), + ) + debit = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .where(col(ApiKey.balance) == key.balance) + .where(col(ApiKey.reserved_balance) == key.reserved_balance) + .values(balance=0, reserved_balance=0, reserved_at=None) + ) + try: + debited = await session.exec(debit) # type: ignore[call-overload] + if debited.rowcount == 0: + await session.rollback() + raise HTTPException( + status_code=409, + detail="Balance changed concurrently. Please retry the refund.", + ) + session.add(refund) + await session.commit() + except IntegrityError: + await session.rollback() + raise HTTPException( + status_code=409, + detail={ + "error": { + "message": "A refund for this key is already in progress.", + "type": "invalid_request_error", + "code": "refund_in_progress", + } + }, + ) + return refund + + +async def _close(session: AsyncSession, refund: Refund, **values: object) -> bool: + result = await session.exec( # type: ignore[call-overload] + update(Refund) + .where(col(Refund.id) == refund.id) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .values(claimed_at=None, updated_at=int(time.time()), **values) + ) + return bool(result.rowcount) + + +async def record_quote(refund: Refund, quote_id: str) -> None: + """Persist the melt quote before the melt is dispatched. + + Once the quote is on disk the reconciler can ask the mint what became of + it, so a crash after this point can never be mistaken for "never sent". + Raises if the claim is no longer open, which aborts the payout. + """ + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(Refund) + .where(col(Refund.id) == refund.id) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .values(quote_id=quote_id, updated_at=int(time.time())) + ) + await session.commit() + if not result.rowcount: + raise LNURLError("Refund claim closed before the melt was dispatched") + refund.quote_id = quote_id + + +async def settle( + session: AsyncSession, + refund: Refund, + *, + quote_id: str | None = None, + token: str | None = None, + mint_url: str | None = None, +) -> bool: + values: dict[str, object] = {"status": "paid"} + if quote_id is not None: + values["quote_id"] = quote_id + if token is not None: + values["token"] = token + if mint_url is not None: + values["mint_url"] = mint_url + settled = await _close(session, refund, **values) + await session.commit() + if not settled: + logger.warning( + "refund paid but its claim was already closed", + extra={"refund_id": refund.id, "prior_status": refund.status}, + ) + return settled + + +async def release(session: AsyncSession, refund: Refund) -> bool: + """Close the claim and return the debited balance in the same transaction.""" + if not await _close(session, refund, status="failed"): + # Nothing changed; commit rather than roll back so the session stays + # usable (an async rollback after an ORM-enabled UPDATE expires the + # identity map and later loads fail outside the greenlet). + await session.commit() + return False + await session.exec( # type: ignore[call-overload] + update(ApiKey) + .where(col(ApiKey.hashed_key) == refund.api_key_hashed_key) + .values(balance=col(ApiKey.balance) + refund.amount_msats) + ) + await session.commit() + logger.info( + "refund released; balance restored", + extra={ + "refund_id": refund.id, + "key_hash": refund.api_key_hashed_key[:8], + "restored_msats": refund.amount_msats, + }, + ) + return True + + +async def hold(session: AsyncSession, refund: Refund, quote_id: str | None) -> None: + """Keep the debit and the claim open until reconciliation resolves it.""" + await _close(session, refund, status="ambiguous", quote_id=quote_id) + await session.commit() + + +async def latest_terminal(session: AsyncSession, key: ApiKey) -> Refund | None: + """Most recent paid Lightning refund, for idempotent re-requests. + + Cashu payouts are deliberately excluded: their token lives in + ``cashu_transactions`` whose ``collected``/``swept`` flags decide whether + it may still be handed out. + """ + result = await session.exec( + select(Refund) + .where(Refund.api_key_hashed_key == key.hashed_key) + .where(Refund.status == "paid") + .where(Refund.method == "lightning") + .order_by(col(Refund.created_at).desc()) + ) + return result.first() + + +def describe(refund: Refund) -> dict[str, str]: + body: dict[str, str] = {"refund_id": refund.id, "status": refund.status} + if refund.token: + body["token"] = refund.token + if refund.destination: + body["recipient"] = refund.destination + if refund.unit == "sat": + body["sats"] = str(refund.amount_msats // 1000) + else: + body["msats"] = str(refund.amount_msats) + return body + + +async def execute(session: AsyncSession, refund: Refund) -> dict[str, str]: + """Pay out an open claim, closing it on every outcome the mint makes known.""" + amount = amount_in_unit(refund.amount_msats, refund.unit) + quote_id: str | None = None + + async def capture_quote(quote: str) -> None: + nonlocal quote_id + quote_id = quote + await record_quote(refund, quote) + + try: + if refund.method == "lightning": + await send_to_lnurl( + amount, + refund.unit, + refund.mint_url, + str(refund.destination), + on_melt_quote=capture_quote, + ) + await settle(session, refund, quote_id=quote_id) + else: + token = await send_token(amount, refund.unit, refund.mint_url) + mint_url = token_mint_url(token, refund.mint_url) + await settle(session, refund, token=token, mint_url=mint_url) + await store_cashu_transaction( + token=token, + amount=amount, + unit=refund.unit, + mint_url=mint_url, + typ="out", + collected=False, + source="apikey", + api_key_hashed_key=refund.api_key_hashed_key, + ) + refund.token = token + refund.mint_url = mint_url + except MeltOutcomeAmbiguousError as e: + await hold(session, refund, quote_id) + logger.error( + "refund outcome ambiguous; balance withheld pending reconciliation", + extra={ + "refund_id": refund.id, + "error": str(e), + "key_hash": refund.api_key_hashed_key[:8], + "quote_id": quote_id, + }, + ) + raise HTTPException( + status_code=502, + detail=( + "Refund was dispatched but its outcome is unconfirmed; the " + "balance is withheld until reconciliation completes" + ), + ) + except HTTPException: + await release(session, refund) + raise + except Exception as e: + await release(session, refund) + logger.error( + "refund payout failed", + extra={ + "refund_id": refund.id, + "error": str(e), + "error_type": type(e).__name__, + "key_hash": refund.api_key_hashed_key[:8], + "method": refund.method, + "mint_url": refund.mint_url, + }, + ) + if is_mint_connection_error(e): + raise HTTPException(status_code=503, detail="Mint service unavailable") + raise HTTPException(status_code=500, detail="Refund failed") + + refund.status = "paid" + refund.claimed_at = None + logger.info( + "refund paid", + extra={ + "refund_id": refund.id, + "method": refund.method, + "amount_msats": refund.amount_msats, + "key_hash": refund.api_key_hashed_key[:8], + }, + ) + return describe(refund) + + +async def _lease(refund_id: str, now: int, lease_cutoff: int) -> bool: + async with create_session() as session: + result = await session.exec( # type: ignore[call-overload] + update(Refund) + .where(col(Refund.id) == refund_id) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .where( + col(Refund.claimed_at).is_(None) + | (col(Refund.claimed_at) < lease_cutoff) + ) + .values(claimed_at=now) + ) + await session.commit() + return bool(result.rowcount) + + +async def _reconcile(refund: Refund) -> None: + if refund.method != "lightning": + # A cashu payout leaves no quote to query: the token either reached the + # client or was lost with the process. Close the claim as ``stuck`` so + # the balance stays withheld and the operator is told exactly once. + async with create_session() as session: + if await _close(session, refund, status="stuck"): + await session.commit() + logger.critical( + "cashu refund stuck; balance withheld, manual reconciliation required", + extra={ + "refund_id": refund.id, + "key_hash": refund.api_key_hashed_key[:8], + "amount_msats": refund.amount_msats, + }, + ) + return + + if refund.quote_id is None: + # No melt quote exists, so the mint was never asked to pay. + async with create_session() as session: + await release(session, refund) + return + + status = await check_bolt11_payment_status( + refund.mint_url, refund.unit, refund.quote_id + ) + if status == "paid": + async with create_session() as session: + await settle(session, refund) + elif status == "unpaid": + async with create_session() as session: + await release(session, refund) + else: + logger.warning( + "refund still unresolved at the mint", + extra={"refund_id": refund.id, "melt_status": status}, + ) + + +async def reconcile_once() -> None: + """Resolve open claims whose lease has lapsed. + + A fresh claim is leased to the request paying it out; ``hold`` drops that + lease so an ambiguous outcome is queried on the next pass. + """ + now = int(time.time()) + lease_cutoff = now - settings.refund_claim_timeout_seconds + async with create_session() as session: + result = await session.exec( + select(Refund) + .where(col(Refund.status).in_(REFUND_OPEN_STATUSES)) + .where( + col(Refund.claimed_at).is_(None) + | (col(Refund.claimed_at) < lease_cutoff) + ) + .order_by(col(Refund.created_at)) + .limit(RECONCILE_BATCH_LIMIT) + ) + stale = list(result.all()) + + for refund in stale: + if not await _lease(refund.id, now, lease_cutoff): + continue + try: + await _reconcile(refund) + except Exception as e: + logger.error( + "refund reconciliation failed", + extra={ + "refund_id": refund.id, + "error": str(e), + "error_type": type(e).__name__, + }, + ) + + +async def periodic_refund_reconcile() -> None: + while True: + await asyncio.sleep(settings.refund_reconcile_interval_seconds) + try: + await reconcile_once() + except asyncio.CancelledError: + raise + except Exception as e: + logger.error( + "refund reconcile loop error", + extra={"error": str(e), "error_type": type(e).__name__}, + ) diff --git a/routstr/wallet.py b/routstr/wallet.py index 910d77e3..023cb404 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -8,7 +8,7 @@ from contextlib import asynccontextmanager from contextvars import ContextVar from dataclasses import dataclass from pathlib import Path -from typing import AsyncGenerator, TypedDict +from typing import AsyncGenerator, Awaitable, Callable, TypedDict from urllib.parse import urlsplit, urlunsplit import httpx @@ -1995,7 +1995,14 @@ async def periodic_routstr_fee_payout() -> None: ) -async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: +async def send_to_lnurl( + amount: int, + unit: str, + mint: str, + address: str, + *, + on_melt_quote: Callable[[str], Awaitable[None]] | None = None, +) -> int: async with wallet_operation_guard(): mint = await find_trusted_mint_with_funds(amount, unit, mint, force_reload=True) wallet = await get_wallet(mint, unit) @@ -2003,7 +2010,14 @@ async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int: # Hand over unreserved proofs: raw_send_to_lnurl reserves only once the # destination, the invoice amount and the melt quote have all been # accepted, so a rejected refund cannot strand locked proofs. - return await raw_send_to_lnurl(wallet, available, address, unit, amount=amount) + return await raw_send_to_lnurl( + wallet, + available, + address, + unit, + amount=amount, + on_melt_quote=on_melt_quote, + ) # class Payment: diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index f369182c..f8babb17 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -554,8 +554,8 @@ async def integration_app( patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl), patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token), patch("routstr.wallet.get_balance", testmint_wallet.get_balance), - patch("routstr.balance.send_token", testmint_wallet.send_token), - patch("routstr.balance.send_to_lnurl", testmint_wallet.send_to_lnurl), + patch("routstr.refund.send_token", testmint_wallet.send_token), + patch("routstr.refund.send_to_lnurl", testmint_wallet.send_to_lnurl), patch("websockets.connect") as mock_websockets, patch("routstr.payment.price.btc_usd_price", return_value=50000.0), patch("routstr.payment.price.sats_usd_price", return_value=0.0005), diff --git a/tests/integration/test_error_handling_edge_cases.py b/tests/integration/test_error_handling_edge_cases.py index 4587fd69..cfd8c221 100644 --- a/tests/integration/test_error_handling_edge_cases.py +++ b/tests/integration/test_error_handling_edge_cases.py @@ -31,7 +31,7 @@ class TestNetworkFailureScenarios: AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), patch( - "routstr.balance.send_token", + "routstr.refund.send_token", AsyncMock(side_effect=ConnectError("Mint service unavailable")), ), ): diff --git a/tests/integration/test_refund_claims.py b/tests/integration/test_refund_claims.py new file mode 100644 index 00000000..0ba4d7fa --- /dev/null +++ b/tests/integration/test_refund_claims.py @@ -0,0 +1,527 @@ +"""Refund claim lifecycle against a real SQLite database. + +Covers the guarantees the ``refunds`` table exists to provide: one open claim +per key, a persisted melt quote before the melt is dispatched, and a +reconciler that never restores a balance whose payout may have settled. +""" + +import time +from typing import Any, Awaitable, Callable +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import HTTPException +from sqlmodel import select + +from routstr import refund +from routstr.balance import RefundRequest, refund_wallet_endpoint +from routstr.core.db import ApiKey, AsyncSession, Refund +from routstr.payment.lnurl import LNURLError, MeltOutcomeAmbiguousError + +KEY_HASH = "refundclaimkey" +ADDRESS = "user@ln.example.com" +BALANCE_MSATS = 5_000_000 + + +async def _seed_key( + session: AsyncSession, *, balance: int = BALANCE_MSATS, address: str | None = None +) -> ApiKey: + key = ApiKey(hashed_key=KEY_HASH) + key.balance = balance + key.reserved_balance = 0 + key.refund_currency = "sat" + key.refund_address = address + key.total_spent = 0 + key.total_requests = 0 + session.add(key) + await session.commit() + await session.refresh(key) + return key + + +async def _load_key(session: AsyncSession) -> ApiKey: + key = await session.get(ApiKey, KEY_HASH) + assert key is not None + await session.refresh(key) + return key + + +async def _load_refund(session: AsyncSession, refund_id: str) -> Refund: + row = await session.get(Refund, refund_id) + assert row is not None + await session.refresh(row) + return row + + +async def _age_claim(session: AsyncSession, refund_id: str, seconds: int) -> None: + row = await _load_refund(session, refund_id) + row.claimed_at = int(time.time()) - seconds + row.created_at = int(time.time()) - seconds + session.add(row) + await session.commit() + + +@pytest.fixture +def short_timeout() -> Any: + with patch.object(refund.settings, "refund_claim_timeout_seconds", 300): + yield + + +# --- exclusivity ----------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_second_claim_on_open_key_is_rejected( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + first = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + key = await _load_key(integration_session) + assert key.balance == 0 + + with pytest.raises(HTTPException) as exc_info: + await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + detail = exc_info.value.detail + assert exc_info.value.status_code == 409 + assert isinstance(detail, dict) + assert detail["error"]["code"] == "refund_in_progress" + + rows = (await integration_session.exec(select(Refund))).all() + assert [row.id for row in rows] == [first.id] + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_claim_from_second_session_hits_the_index( + integration_engine: Any, integration_session: AsyncSession +) -> None: + key = await _seed_key(integration_session) + await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + + async with AsyncSession(integration_engine, expire_on_commit=False) as other: + other_key = await _load_key(other) + with pytest.raises(HTTPException) as exc_info: + await refund.open_claim( + other, other_key, method="lightning", destination=ADDRESS + ) + assert exc_info.value.status_code == 409 + + +@pytest.mark.asyncio +async def test_claim_rejects_stale_balance_snapshot( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + integration_session.expunge(key) # a stale, detached snapshot + key.balance = BALANCE_MSATS + 1 + with pytest.raises(HTTPException) as exc_info: + await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + assert exc_info.value.status_code == 409 + assert (await integration_session.exec(select(Refund))).all() == [] + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_retry_after_failed_claim_pays_once( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + first = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + assert await refund.release(integration_session, first) + key = await _load_key(integration_session) + assert key.balance == BALANCE_MSATS + + second = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + assert second.id != first.id + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_release_after_settle_is_a_noop( + integration_session: AsyncSession, +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + assert await refund.settle(integration_session, claim, quote_id="q1") + assert not await refund.release(integration_session, claim) + assert (await _load_key(integration_session)).balance == 0 + row = await _load_refund(integration_session, claim.id) + assert row.status == "paid" + assert row.claimed_at is None + + +# --- execute --------------------------------------------------------------- + + +def _lnurl_stub( + outcome: BaseException | None = None, +) -> Callable[..., Awaitable[int]]: + async def send( + amount: int, + unit: str, + mint: str, + address: str, + *, + on_melt_quote: Callable[[str], Awaitable[None]] | None = None, + ) -> int: + if on_melt_quote is not None: + await on_melt_quote("quote-123") + if outcome is not None: + raise outcome + return amount + + return send + + +@pytest.mark.asyncio +async def test_execute_persists_quote_before_melt_and_settles( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + seen: list[str | None] = [] + + async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int: + await on_melt_quote("quote-123") + row = await _load_refund(integration_session, claim.id) + seen.append(row.quote_id) + return 5000 + + with patch("routstr.refund.send_to_lnurl", send): + body = await refund.execute(integration_session, claim) + + assert seen == ["quote-123"], "quote must be on disk before the melt runs" + assert body["status"] == "paid" + assert body["recipient"] == ADDRESS + assert body["sats"] == "5000" + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.quote_id, row.claimed_at) == ("paid", "quote-123", None) + + +@pytest.mark.asyncio +async def test_execute_ambiguous_holds_claim_with_quote( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch( + "routstr.refund.send_to_lnurl", _lnurl_stub(MeltOutcomeAmbiguousError("?")) + ): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + + assert exc_info.value.status_code == 502 + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.quote_id, row.claimed_at) == ( + "ambiguous", + "quote-123", + None, + ) + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_execute_clean_failure_restores_balance( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch("routstr.refund.send_to_lnurl", _lnurl_stub(LNURLError("limits"))): + with pytest.raises(HTTPException) as exc_info: + await refund.execute(integration_session, claim) + + assert exc_info.value.status_code == 500 + row = await _load_refund(integration_session, claim.id) + assert row.status == "failed" + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_execute_aborts_melt_when_claim_was_released( + integration_engine: Any, integration_session: AsyncSession, patched_db_engine: None +) -> None: + """A reconciler that released the claim first must stop the melt.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + melted = False + + async def send(*args: Any, on_melt_quote: Any = None, **kwargs: Any) -> int: + nonlocal melted + async with AsyncSession(integration_engine, expire_on_commit=False) as other: + await refund.release(other, await _load_refund(other, claim.id)) + await on_melt_quote("quote-123") + melted = True + return 5000 + + with patch("routstr.refund.send_to_lnurl", send): + with pytest.raises(HTTPException): + await refund.execute(integration_session, claim) + + assert melted is False + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +# --- reconciler ------------------------------------------------------------ + + +async def _open_ambiguous(session: AsyncSession, quote_id: str | None) -> Refund: + key = await _seed_key(session) + claim = await refund.open_claim( + session, key, method="lightning", destination=ADDRESS + ) + await refund.hold(session, claim, quote_id) + return claim + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("mint_status", "expected_status", "expected_balance"), + [ + ("paid", "paid", 0), + ("unpaid", "failed", BALANCE_MSATS), + ("pending", "ambiguous", 0), + ("unknown", "ambiguous", 0), + ], +) +async def test_reconcile_ambiguous_claims( + integration_session: AsyncSession, + patched_db_engine: None, + short_timeout: None, + mint_status: str, + expected_status: str, + expected_balance: int, +) -> None: + claim = await _open_ambiguous(integration_session, "quote-123") + with patch( + "routstr.refund.check_bolt11_payment_status", + AsyncMock(return_value=mint_status), + ) as check: + await refund.reconcile_once() + + check.assert_awaited_once_with(claim.mint_url, "sat", "quote-123") + row = await _load_refund(integration_session, claim.id) + assert row.status == expected_status + assert (await _load_key(integration_session)).balance == expected_balance + + +@pytest.mark.asyncio +async def test_reconcile_credits_balance_once_across_passes( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + await _open_ambiguous(integration_session, "quote-123") + with patch( + "routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="unpaid") + ): + await refund.reconcile_once() + await refund.reconcile_once() + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_reconcile_leaves_fresh_pending_claim_alone( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check: + await refund.reconcile_once() + check.assert_not_awaited() + row = await _load_refund(integration_session, claim.id) + assert row.status == "pending" + assert row.claimed_at is not None + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_reconcile_releases_expired_claim_without_quote( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + """No quote on disk means the mint was never asked to pay.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + await _age_claim(integration_session, claim.id, 600) + with patch("routstr.refund.check_bolt11_payment_status", AsyncMock()) as check: + await refund.reconcile_once() + check.assert_not_awaited() + assert (await _load_refund(integration_session, claim.id)).status == "failed" + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_reconcile_queries_mint_for_crashed_claim_with_quote( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + """Crash after the quote was persisted: the mint decides, not the timeout.""" + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="lightning", destination=ADDRESS + ) + await refund.record_quote(claim, "quote-crash") + await _age_claim(integration_session, claim.id, 600) + with patch( + "routstr.refund.check_bolt11_payment_status", AsyncMock(return_value="paid") + ) as check: + await refund.reconcile_once() + check.assert_awaited_once_with(claim.mint_url, "sat", "quote-crash") + assert (await _load_refund(integration_session, claim.id)).status == "paid" + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_reconcile_marks_expired_cashu_claim_stuck( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + key = await _seed_key(integration_session) + claim = await refund.open_claim( + integration_session, key, method="cashu", destination=None + ) + await _age_claim(integration_session, claim.id, 600) + with patch("routstr.refund.logger") as log: + await refund.reconcile_once() + await refund.reconcile_once() + assert log.critical.call_count == 1 + row = await _load_refund(integration_session, claim.id) + assert (row.status, row.claimed_at) == ("stuck", None) + assert (await _load_key(integration_session)).balance == 0 + # A stuck claim is closed, so the key is not permanently locked out. + key = await _load_key(integration_session) + key.balance = 1000 + integration_session.add(key) + await integration_session.commit() + await refund.open_claim(integration_session, key, method="cashu", destination=None) + + +@pytest.mark.asyncio +async def test_reconcile_survives_one_failing_row( + integration_session: AsyncSession, patched_db_engine: None, short_timeout: None +) -> None: + claim = await _open_ambiguous(integration_session, "quote-123") + with patch( + "routstr.refund.check_bolt11_payment_status", + AsyncMock(side_effect=RuntimeError("mint down")), + ): + await refund.reconcile_once() + row = await _load_refund(integration_session, claim.id) + assert row.status == "ambiguous" + assert row.claimed_at is not None, "lease is kept until the next pass" + + +# --- endpoint -------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_endpoint_uses_requested_address_over_persisted( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _seed_key(integration_session, address="stored@ln.example.com") + send = AsyncMock(side_effect=_lnurl_stub()) + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()) as resolve, + patch("routstr.refund.send_to_lnurl", send), + ): + body = await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address=ADDRESS), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + resolve.assert_awaited_once_with(ADDRESS) + assert isinstance(body, dict) + assert body["recipient"] == ADDRESS + assert send.await_args is not None + assert send.await_args.args[3] == ADDRESS + assert (await _load_key(integration_session)).balance == 0 + + +@pytest.mark.asyncio +async def test_endpoint_rejects_bad_address_without_debit( + integration_session: AsyncSession, +) -> None: + await _seed_key(integration_session) + with patch( + "routstr.refund.get_lnurl_data", AsyncMock(side_effect=LNURLError("nope")) + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address="bad@example"), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 400 + assert (await integration_session.exec(select(Refund))).all() == [] + assert (await _load_key(integration_session)).balance == BALANCE_MSATS + + +@pytest.mark.asyncio +async def test_endpoint_replays_paid_lightning_refund_on_empty_balance( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _seed_key(integration_session, address=ADDRESS) + with patch("routstr.refund.send_to_lnurl", _lnurl_stub()): + first = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + second = await refund_wallet_endpoint( + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert isinstance(first, dict) and isinstance(second, dict) + assert second["refund_id"] == first["refund_id"] + assert second["status"] == "paid" + + +@pytest.mark.asyncio +async def test_endpoint_refund_while_ambiguous_returns_409( + integration_session: AsyncSession, patched_db_engine: None +) -> None: + await _open_ambiguous(integration_session, "quote-123") + key = await _load_key(integration_session) + key.balance = 2_000_000 # topped up while the melt is unresolved + integration_session.add(key) + await integration_session.commit() + + send = AsyncMock() + with ( + patch("routstr.refund.get_lnurl_data", AsyncMock()), + patch("routstr.refund.send_to_lnurl", send), + ): + with pytest.raises(HTTPException) as exc_info: + await refund_wallet_endpoint( + refund_request=RefundRequest(lightning_address=ADDRESS), + authorization=f"Bearer sk-{KEY_HASH}", + x_cashu=None, + session=integration_session, + ) + assert exc_info.value.status_code == 409 + send.assert_not_awaited() + assert (await _load_key(integration_session)).balance == 2_000_000 diff --git a/tests/integration/test_wallet_refund.py b/tests/integration/test_wallet_refund.py index 1cbd6b4b..83f4ec7f 100644 --- a/tests/integration/test_wallet_refund.py +++ b/tests/integration/test_wallet_refund.py @@ -204,7 +204,7 @@ async def test_refund_with_lightning_address( await db_snapshot.capture() # Mock send_to_lnurl function directly - with patch("routstr.balance.send_to_lnurl") as mock_send_to_lnurl: + with patch("routstr.refund.send_to_lnurl") as mock_send_to_lnurl: mock_send_to_lnurl.return_value = { "amount_sent": balance, "unit": "msat", @@ -508,7 +508,7 @@ async def test_mint_unavailability_handling( # Make the send_token method raise a typed mint connection exception. with patch( - "routstr.balance.send_token", + "routstr.refund.send_token", side_effect=MintConnectionError(raw_error), ): response = await authenticated_client.post("/v1/wallet/refund") @@ -622,7 +622,7 @@ async def test_refund_with_expired_key( integration_client.headers["Authorization"] = f"Bearer {api_key}" # Mock the refund to LN address - with patch("routstr.balance.send_to_lnurl") as mock_send_to_lnurl: + with patch("routstr.refund.send_to_lnurl") as mock_send_to_lnurl: mock_send_to_lnurl.return_value = 500 response = await integration_client.post("/v1/wallet/refund") diff --git a/tests/unit/test_balance.py b/tests/unit/test_balance.py index 75045168..cc970d40 100644 --- a/tests/unit/test_balance.py +++ b/tests/unit/test_balance.py @@ -20,7 +20,9 @@ def _make_cashu_tx( swept: bool = False, collected: bool = False, ) -> CashuTransaction: - tx = CashuTransaction(token=token, amount=amount, unit=unit, type=type, request_id=request_id) + tx = CashuTransaction( + token=token, amount=amount, unit=unit, type=type, request_id=request_id + ) tx.swept = swept tx.collected = collected return tx @@ -41,13 +43,22 @@ def _update_result(rowcount: int) -> MagicMock: @pytest.mark.asyncio async def test_refund_x_cashu_returns_token() -> None: x_cashu_token = "cashuAtest_token_value" - in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc") - out_tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat", type="out", request_id="req-abc") + in_tx = _make_cashu_tx( + token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc" + ) + out_tx = _make_cashu_tx( + token="cashuArefund_token", + amount=1000, + unit="msat", + type="out", + request_id="req-abc", + ) session = MagicMock() session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() result = await refund_wallet_endpoint( authorization="Bearer sk-somekey", @@ -66,13 +77,22 @@ async def test_refund_x_cashu_returns_token() -> None: @pytest.mark.asyncio async def test_refund_x_cashu_sat_unit() -> None: x_cashu_token = "cashuAsat_token" - in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat") - out_tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat", type="out", request_id="req-sat") + in_tx = _make_cashu_tx( + token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat" + ) + out_tx = _make_cashu_tx( + token="cashuArefund_sat", + amount=500, + unit="sat", + type="out", + request_id="req-sat", + ) session = MagicMock() session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() result = await refund_wallet_endpoint( authorization="Bearer sk-somekey", @@ -124,6 +144,7 @@ async def test_refund_x_cashu_pending_raises_425() -> None: session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(None)]) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -167,8 +188,21 @@ async def test_refund_x_cashu_in_tx_without_request_id_raises_404() -> None: async def test_refund_x_cashu_swept_raises_410() -> None: from fastapi import HTTPException - in_tx = _make_cashu_tx(token="cashuAswept_token", amount=0, unit="msat", type="in", request_id="req-swept") - out_tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", type="out", request_id="req-swept", swept=True) + in_tx = _make_cashu_tx( + token="cashuAswept_token", + amount=0, + unit="msat", + type="in", + request_id="req-swept", + ) + out_tx = _make_cashu_tx( + token="cashuAswept", + amount=100, + unit="msat", + type="out", + request_id="req-swept", + swept=True, + ) session = MagicMock() session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) @@ -239,10 +273,11 @@ async def test_apikey_refund_returns_persisted_token_after_cache_loss() -> None: session.exec = AsyncMock(return_value=_exec_result(refund_tx)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance.send_token", AsyncMock()) as mock_send_token, + patch("routstr.refund.latest_terminal", AsyncMock(return_value=None)), + patch("routstr.refund.send_token", AsyncMock()) as mock_send_token, ): result = await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -277,8 +312,9 @@ async def test_apikey_refund_rejects_persisted_token_after_sweep() -> None: session.exec = AsyncMock(return_value=_exec_result(refund_tx)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() - with patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)): + with patch("routstr.refund.latest_terminal", AsyncMock(return_value=None)): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -302,12 +338,11 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store, - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)), + patch("routstr.refund.store_cashu_transaction", AsyncMock()) as mock_store, ): result = await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -336,13 +371,12 @@ async def test_apikey_refund_logs_token() -> None: session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger") as mock_logger, + patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger") as mock_logger, ): await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -351,11 +385,11 @@ async def test_apikey_refund_logs_token() -> None: ) calls = [str(c) for c in mock_logger.info.call_args_list] - assert any("cashu token issued" in c for c in calls) + assert any("refund paid" in c for c in calls) @pytest.mark.asyncio -async def test_apikey_refund_log_includes_path() -> None: +async def test_apikey_refund_log_identifies_the_claim() -> None: key = _make_api_key(balance=5000, refund_currency="sat") refund_token = "cashuApath_token" @@ -364,13 +398,12 @@ async def test_apikey_refund_log_includes_path() -> None: session.exec = AsyncMock(return_value=_update_result(1)) session.add = MagicMock() session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger") as mock_logger, + patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger") as mock_logger, ): await refund_wallet_endpoint( authorization="Bearer sk-testhash", @@ -378,14 +411,15 @@ async def test_apikey_refund_log_includes_path() -> None: session=session, ) - # Find the "cashu token issued" call and verify extra contains the path - token_issued_calls = [ - c for c in mock_logger.info.call_args_list - if c.args and "cashu token issued" in c.args[0] + paid_calls = [ + c + for c in mock_logger.info.call_args_list + if c.args and "refund paid" in c.args[0] ] - assert len(token_issued_calls) == 1 - extra = token_issued_calls[0].kwargs.get("extra", {}) - assert extra.get("path") == "/v1/wallet/refund" + assert len(paid_calls) == 1 + extra = paid_calls[0].kwargs.get("extra", {}) + assert extra.get("method") == "cashu" + assert extra.get("refund_id") @pytest.mark.asyncio @@ -400,14 +434,13 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None: # Debit returns rowcount=0 → balance changed concurrently session.exec = AsyncMock(return_value=_update_result(0)) session.commit = AsyncMock() + session.rollback = AsyncMock() mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted") with ( - patch("routstr.balance.send_token", mock_send_token), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.refund.send_token", mock_send_token), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -427,6 +460,7 @@ async def test_credit_balance_stores_apikey_transaction_history() -> None: session = MagicMock() session.exec = AsyncMock(return_value=_update_result(1)) session.commit = AsyncMock() + session.rollback = AsyncMock() session.refresh = AsyncMock() with ( @@ -460,18 +494,20 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: # First exec call = debit (succeeds), second = restore session = MagicMock() session.get = AsyncMock(return_value=key) - session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)]) + # debit, then the claim close and the balance restore + session.exec = AsyncMock( + side_effect=[_update_result(1), _update_result(1), _update_result(1)] + ) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( patch( - "routstr.balance.send_token", + "routstr.refund.send_token", AsyncMock(side_effect=MintConnectionError("raw mint outage detail")), ), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger"), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger"), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -483,8 +519,8 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None: assert exc_info.value.status_code == 503 assert exc_info.value.detail == "Mint service unavailable" assert "raw mint outage detail" not in exc_info.value.detail - # Verify two exec calls: debit + restore - assert session.exec.await_count == 2 + # debit, claim close, balance restore + assert session.exec.await_count == 3 @pytest.mark.asyncio @@ -497,15 +533,19 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None: session = MagicMock() session.get = AsyncMock(return_value=key) - session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)]) + # debit, then the claim close and the balance restore + session.exec = AsyncMock( + side_effect=[_update_result(1), _update_result(1), _update_result(1)] + ) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), - patch("routstr.balance.logger"), + patch( + "routstr.refund.send_token", AsyncMock(side_effect=RuntimeError(raw_error)) + ), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), + patch("routstr.refund.logger"), ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -517,7 +557,7 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None: assert exc_info.value.status_code == 500 assert exc_info.value.detail == "Refund failed" assert raw_error not in exc_info.value.detail - assert session.exec.await_count == 2 + assert session.exec.await_count == 3 # --------------------------------------------------------------------------- @@ -575,6 +615,7 @@ async def test_refund_unknown_sk_bearer_returns_401() -> None: # --- Topup redemption error taxonomy (POST /v1/wallet/topup) ------------------ + def _envelope(exc: HTTPException) -> dict: """Extract the error object from a top-up HTTPException.""" detail = exc.detail @@ -633,7 +674,9 @@ async def test_topup_mint_unreachable_returns_503( @pytest.mark.asyncio -async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> None: +async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> ( + None +): from fastapi import HTTPException from routstr.wallet import SourceMintConnectionError @@ -697,7 +740,9 @@ async def test_topup_zero_value_returns_400_zero_value_message() -> None: patch( "routstr.balance.credit_balance", AsyncMock( - side_effect=ValueError("Redeemed token amount must be positive, got 0 msats") + side_effect=ValueError( + "Redeemed token amount must be positive, got 0 msats" + ) ), ), ): @@ -758,9 +803,7 @@ async def test_topup_token_consumed_returns_500() -> None: "Token value is too small to cover swap fees", ), ( - ValueError( - "Token amount (5 sat) is insufficient to cover melt fees." - ), + ValueError("Token amount (5 sat) is insufficient to cover melt fees."), 422, "mint_error", "cashu_token_swap_fees_exceed_amount", @@ -837,15 +880,14 @@ async def test_apikey_refund_ambiguous_melt_does_not_restore_balance() -> None: session.get = AsyncMock(return_value=key) session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), patch( - "routstr.balance.send_to_lnurl", + "routstr.refund.send_to_lnurl", AsyncMock(side_effect=MeltOutcomeAmbiguousError("outcome is ambiguous")), ), - patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore, + patch("routstr.refund.release", AsyncMock()) as mock_restore, ): with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( @@ -869,15 +911,14 @@ async def test_apikey_refund_clean_failure_still_restores_balance() -> None: session.get = AsyncMock(return_value=key) session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) session.commit = AsyncMock() + session.rollback = AsyncMock() with ( - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), patch( - "routstr.balance.send_to_lnurl", + "routstr.refund.send_to_lnurl", AsyncMock(side_effect=RuntimeError("mint rejected melt")), ), - patch("routstr.balance._restore_balance", AsyncMock()) as mock_restore, + patch("routstr.refund.release", AsyncMock()) as mock_restore, ): with pytest.raises(HTTPException): await refund_wallet_endpoint( diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 47061dda..498b57f5 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -200,10 +200,8 @@ async def test_reset_all_reserved_balances_clears_reserved_at( def _refund_patches(refund_token: str = "cashuArefund"): # type: ignore[no-untyped-def] return ( - patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), - patch("routstr.balance.store_cashu_transaction", AsyncMock()), - patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), - patch("routstr.balance._refund_cache_set", AsyncMock()), + patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)), + patch("routstr.refund.store_cashu_transaction", AsyncMock()), ) @@ -224,8 +222,8 @@ async def test_refund_self_heals_stale_reservation(session: AsyncSession) -> Non reserved_at=int(time.time()) - 10_000, ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: result = await refund_wallet_endpoint( authorization="Bearer sk-stalerefund", x_cashu=None, @@ -253,8 +251,8 @@ async def test_refund_self_heals_legacy_null_reserved_at(session: AsyncSession) reserved_at=None, ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: result = await refund_wallet_endpoint( authorization="Bearer sk-legacyrefund", x_cashu=None, @@ -280,8 +278,8 @@ async def test_refund_rejects_recent_reservation(session: AsyncSession) -> None: reserved_at=int(time.time()), ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: with pytest.raises(HTTPException) as exc_info: await refund_wallet_endpoint( authorization="Bearer sk-activerefund", @@ -302,8 +300,8 @@ async def test_refund_without_reservation_still_works(session: AsyncSession) -> reserved_balance=0, ) - p1, p2, p3, p4 = _refund_patches() - with p1, p2, p3, p4: + p1, p2 = _refund_patches() + with p1, p2: result = await refund_wallet_endpoint( authorization="Bearer sk-plainrefund", x_cashu=None,