added lightning refund

This commit is contained in:
9qeklajc
2026-09-12 02:20:09 +02:00
parent efe9c5599d
commit eec227a88f
13 changed files with 1240 additions and 307 deletions
@@ -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")
+21 -216
View File
@@ -1,13 +1,12 @@
import asyncio
import hashlib import hashlib
from time import monotonic
from typing import Annotated, NoReturn from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException from fastapi import APIRouter, Depends, Header, HTTPException
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from pydantic import BaseModel from pydantic import BaseModel
from sqlmodel import col, select, update from sqlmodel import col, select
from . import refund
from .auth import ( from .auth import (
redemption_error_to_http_exception, redemption_error_to_http_exception,
validate_bearer_key, validate_bearer_key,
@@ -19,20 +18,13 @@ from .core.db import (
get_session, get_session,
release_stale_reservations, release_stale_reservations,
) )
from .core.db import (
store_cashu_transaction_with_retry as store_cashu_transaction,
)
from .core.logging import get_logger from .core.logging import get_logger
from .core.settings import settings from .core.settings import settings
from .lightning import lightning_router from .lightning import lightning_router
from .payment.lnurl import MeltOutcomeAmbiguousError
from .wallet import ( from .wallet import (
classify_redemption_error, classify_redemption_error,
credit_balance, credit_balance,
is_mint_connection_error,
recieve_token, recieve_token,
send_to_lnurl,
send_token,
token_mint_url, token_mint_url,
) )
@@ -236,35 +228,6 @@ async def topup_wallet_endpoint(
return {"msats": amount_msats} 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( async def _lookup_key_no_create(
bearer_value: str, session: AsyncSession bearer_value: str, session: AsyncSession
) -> ApiKey | None: ) -> ApiKey | None:
@@ -307,36 +270,13 @@ async def _get_persisted_api_key_refund(
return persisted return persisted
async def _restore_balance( class RefundRequest(BaseModel):
session: AsyncSession, lightning_address: str | None = None
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,
},
)
@router.post("/refund", response_model=None) @router.post("/refund", response_model=None)
async def refund_wallet_endpoint( async def refund_wallet_endpoint(
refund_request: RefundRequest | None = None,
authorization: Annotated[str | None, Header()] = None, authorization: Annotated[str | None, Header()] = None,
x_cashu: Annotated[str | None, Header()] = None, x_cashu: Annotated[str | None, Header()] = None,
session: AsyncSession = Depends(get_session), 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 key.total_balance <= 0:
if cached := await _refund_cache_get(bearer_value): if paid := await refund.latest_terminal(session, key):
return cached return refund.describe(paid)
if persisted := await _get_persisted_api_key_refund(key, session): if persisted := await _get_persisted_api_key_refund(key, session):
return persisted return persisted
@@ -441,162 +386,21 @@ async def refund_wallet_endpoint(
) )
remaining_balance_msats: int = key.total_balance remaining_balance_msats: int = key.total_balance
unit = refund.refund_unit(key)
if key.refund_currency == "sat": remaining_balance = refund.amount_in_unit(remaining_balance_msats, unit)
remaining_balance = remaining_balance_msats // 1000
else:
remaining_balance = remaining_balance_msats
if remaining_balance_msats > 0 and remaining_balance <= 0: if remaining_balance_msats > 0 and remaining_balance <= 0:
raise HTTPException(status_code=400, detail="Balance too small to refund") raise HTTPException(status_code=400, detail="Balance too small to refund")
elif remaining_balance <= 0: elif remaining_balance <= 0:
raise HTTPException(status_code=400, detail="No balance to refund") raise HTTPException(status_code=400, detail="No balance to refund")
# Capture values before debit — the session may refresh key after commit claim = await refund.open_claim(
pre_debit_balance = key.balance session,
pre_debit_reserved = key.reserved_balance key,
method="lightning" if destination else "cashu",
# --- DEBIT FIRST: atomically zero the balance before minting tokens --- destination=destination,
# 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)
) )
debit_result = await session.exec(debit_stmt) # type: ignore[call-overload] return await refund.execute(session, claim)
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
@router.get("/history") @router.get("/history")
@@ -640,6 +444,7 @@ async def donate(token: str, ref: str | None = None) -> str:
except Exception: except Exception:
return "Invalid token." return "Invalid token."
@router.api_route( @router.api_route(
"/{path:path}", "/{path:path}",
methods=["GET", "POST", "PUT", "DELETE"], methods=["GET", "POST", "PUT", "DELETE"],
+51 -4
View File
@@ -12,7 +12,7 @@ from typing import AsyncGenerator
from alembic import command from alembic import command
from alembic.config import Config from alembic.config import Config
from alembic.util.exc import CommandError 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.engine import make_url
from sqlalchemy.exc import IntegrityError, OperationalError from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.ext.asyncio import AsyncEngine from sqlalchemy.ext.asyncio import AsyncEngine
@@ -300,9 +300,7 @@ async def release_stale_reservations(
col(ApiKey.reserved_at) < cutoff col(ApiKey.reserved_at) < cutoff
) )
else: else:
legacy_query = legacy_query.where( legacy_query = legacy_query.where(col(ApiKey.hashed_key) == key_hash).where(
col(ApiKey.hashed_key) == key_hash
).where(
or_(col(ApiKey.reserved_at).is_(None), col(ApiKey.reserved_at) < cutoff) 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( async def store_cashu_transaction(
token: str, token: str,
amount: int, amount: int,
+7
View File
@@ -28,6 +28,7 @@ from ..nostr.discovery import providers_router
from ..payment.models import models_router, update_sats_pricing from ..payment.models import models_router, update_sats_pricing
from ..payment.price import update_prices_periodically from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_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.auto_topup import periodic_auto_topup
from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
from ..upstream.litellm_routing import configure_litellm from ..upstream.litellm_routing import configure_litellm
@@ -68,6 +69,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
dead_key_prune_task = None dead_key_prune_task = None
auto_topup_task = None auto_topup_task = None
refund_sweep_task = None refund_sweep_task = None
refund_reconcile_task = None
routstr_fee_task = None routstr_fee_task = None
invoice_watcher_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()) dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune())
auto_topup_task = asyncio.create_task(periodic_auto_topup()) auto_topup_task = asyncio.create_task(periodic_auto_topup())
refund_sweep_task = asyncio.create_task(periodic_refund_sweep()) 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()) routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher()) invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher())
@@ -201,6 +204,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
auto_topup_task.cancel() auto_topup_task.cancel()
if refund_sweep_task is not None: if refund_sweep_task is not None:
refund_sweep_task.cancel() refund_sweep_task.cancel()
if refund_reconcile_task is not None:
refund_reconcile_task.cancel()
if routstr_fee_task is not None: if routstr_fee_task is not None:
routstr_fee_task.cancel() routstr_fee_task.cancel()
if invoice_watcher_task is not None: 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) tasks_to_wait.append(auto_topup_task)
if refund_sweep_task is not None: if refund_sweep_task is not None:
tasks_to_wait.append(refund_sweep_task) 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: if routstr_fee_task is not None:
tasks_to_wait.append(routstr_fee_task) tasks_to_wait.append(routstr_fee_task)
if invoice_watcher_task is not None: if invoice_watcher_task is not None:
+8 -1
View File
@@ -124,7 +124,6 @@ class Settings(BaseSettings):
enable_model_paths_refresh: bool = Field( enable_model_paths_refresh: bool = Field(
default=True, env="ENABLE_MODEL_PATHS_REFRESH" 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). # Uncollected refund tokens are swept after ~6 months (180 days).
# Fixed for now: not configurable via env or the settings DB/admin API # Fixed for now: not configurable via env or the settings DB/admin API
# (empty env list disables env binding; see FIXED_FIELDS). # (empty env list disables env binding; see FIXED_FIELDS).
@@ -132,6 +131,14 @@ class Settings(BaseSettings):
refund_sweep_claim_timeout_seconds: int = Field( refund_sweep_claim_timeout_seconds: int = Field(
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS" 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 # Database connection-pool controls (advanced). Capacity defaults provide
# headroom for Routstr's concurrent request and background-payment workload. # headroom for Routstr's concurrent request and background-payment workload.
+429
View File
@@ -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__},
)
+17 -3
View File
@@ -8,7 +8,7 @@ from contextlib import asynccontextmanager
from contextvars import ContextVar from contextvars import ContextVar
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import AsyncGenerator, TypedDict from typing import AsyncGenerator, Awaitable, Callable, TypedDict
from urllib.parse import urlsplit, urlunsplit from urllib.parse import urlsplit, urlunsplit
import httpx 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(): async with wallet_operation_guard():
mint = await find_trusted_mint_with_funds(amount, unit, mint, force_reload=True) mint = await find_trusted_mint_with_funds(amount, unit, mint, force_reload=True)
wallet = await get_wallet(mint, unit) 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 # Hand over unreserved proofs: raw_send_to_lnurl reserves only once the
# destination, the invoice amount and the melt quote have all been # destination, the invoice amount and the melt quote have all been
# accepted, so a rejected refund cannot strand locked proofs. # 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: # class Payment:
+2 -2
View File
@@ -554,8 +554,8 @@ async def integration_app(
patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl), patch("routstr.wallet.send_to_lnurl", testmint_wallet.send_to_lnurl),
patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token), patch("routstr.wallet.recieve_token", testmint_wallet.redeem_token),
patch("routstr.wallet.get_balance", testmint_wallet.get_balance), patch("routstr.wallet.get_balance", testmint_wallet.get_balance),
patch("routstr.balance.send_token", testmint_wallet.send_token), patch("routstr.refund.send_token", testmint_wallet.send_token),
patch("routstr.balance.send_to_lnurl", testmint_wallet.send_to_lnurl), patch("routstr.refund.send_to_lnurl", testmint_wallet.send_to_lnurl),
patch("websockets.connect") as mock_websockets, patch("websockets.connect") as mock_websockets,
patch("routstr.payment.price.btc_usd_price", return_value=50000.0), patch("routstr.payment.price.btc_usd_price", return_value=50000.0),
patch("routstr.payment.price.sats_usd_price", return_value=0.0005), patch("routstr.payment.price.sats_usd_price", return_value=0.0005),
@@ -31,7 +31,7 @@ class TestNetworkFailureScenarios:
AsyncMock(side_effect=ConnectError("Mint service unavailable")), AsyncMock(side_effect=ConnectError("Mint service unavailable")),
), ),
patch( patch(
"routstr.balance.send_token", "routstr.refund.send_token",
AsyncMock(side_effect=ConnectError("Mint service unavailable")), AsyncMock(side_effect=ConnectError("Mint service unavailable")),
), ),
): ):
+527
View File
@@ -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
+3 -3
View File
@@ -204,7 +204,7 @@ async def test_refund_with_lightning_address(
await db_snapshot.capture() await db_snapshot.capture()
# Mock send_to_lnurl function directly # 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 = { mock_send_to_lnurl.return_value = {
"amount_sent": balance, "amount_sent": balance,
"unit": "msat", "unit": "msat",
@@ -508,7 +508,7 @@ async def test_mint_unavailability_handling(
# Make the send_token method raise a typed mint connection exception. # Make the send_token method raise a typed mint connection exception.
with patch( with patch(
"routstr.balance.send_token", "routstr.refund.send_token",
side_effect=MintConnectionError(raw_error), side_effect=MintConnectionError(raw_error),
): ):
response = await authenticated_client.post("/v1/wallet/refund") 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}" integration_client.headers["Authorization"] = f"Bearer {api_key}"
# Mock the refund to LN address # 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 mock_send_to_lnurl.return_value = 500
response = await integration_client.post("/v1/wallet/refund") response = await integration_client.post("/v1/wallet/refund")
+106 -65
View File
@@ -20,7 +20,9 @@ def _make_cashu_tx(
swept: bool = False, swept: bool = False,
collected: bool = False, collected: bool = False,
) -> CashuTransaction: ) -> 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.swept = swept
tx.collected = collected tx.collected = collected
return tx return tx
@@ -41,13 +43,22 @@ def _update_result(rowcount: int) -> MagicMock:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_refund_x_cashu_returns_token() -> None: async def test_refund_x_cashu_returns_token() -> None:
x_cashu_token = "cashuAtest_token_value" 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") in_tx = _make_cashu_tx(
out_tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat", type="out", request_id="req-abc") 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 = MagicMock()
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)])
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-somekey", authorization="Bearer sk-somekey",
@@ -66,13 +77,22 @@ async def test_refund_x_cashu_returns_token() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_refund_x_cashu_sat_unit() -> None: async def test_refund_x_cashu_sat_unit() -> None:
x_cashu_token = "cashuAsat_token" x_cashu_token = "cashuAsat_token"
in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat") in_tx = _make_cashu_tx(
out_tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat", type="out", request_id="req-sat") 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 = MagicMock()
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)])
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-somekey", 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.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(None)])
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with pytest.raises(HTTPException) as exc_info: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( 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: async def test_refund_x_cashu_swept_raises_410() -> None:
from fastapi import HTTPException from fastapi import HTTPException
in_tx = _make_cashu_tx(token="cashuAswept_token", amount=0, unit="msat", type="in", request_id="req-swept") in_tx = _make_cashu_tx(
out_tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", type="out", request_id="req-swept", swept=True) 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 = MagicMock()
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)]) 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.exec = AsyncMock(return_value=_exec_result(refund_tx))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), patch("routstr.refund.latest_terminal", AsyncMock(return_value=None)),
patch("routstr.balance.send_token", AsyncMock()) as mock_send_token, patch("routstr.refund.send_token", AsyncMock()) as mock_send_token,
): ):
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-testhash", 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.exec = AsyncMock(return_value=_exec_result(refund_tx))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() 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: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( await refund_wallet_endpoint(
authorization="Bearer sk-testhash", 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.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store, patch("routstr.refund.store_cashu_transaction", AsyncMock()) as mock_store,
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
): ):
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-testhash", 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.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.refund.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), patch("routstr.refund.logger") as mock_logger,
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger") as mock_logger,
): ):
await refund_wallet_endpoint( await refund_wallet_endpoint(
authorization="Bearer sk-testhash", 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] 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 @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") key = _make_api_key(balance=5000, refund_currency="sat")
refund_token = "cashuApath_token" 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.exec = AsyncMock(return_value=_update_result(1))
session.add = MagicMock() session.add = MagicMock()
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.refund.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), patch("routstr.refund.logger") as mock_logger,
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger") as mock_logger,
): ):
await refund_wallet_endpoint( await refund_wallet_endpoint(
authorization="Bearer sk-testhash", authorization="Bearer sk-testhash",
@@ -378,14 +411,15 @@ async def test_apikey_refund_log_includes_path() -> None:
session=session, session=session,
) )
# Find the "cashu token issued" call and verify extra contains the path paid_calls = [
token_issued_calls = [ c
c for c in mock_logger.info.call_args_list for c in mock_logger.info.call_args_list
if c.args and "cashu token issued" in c.args[0] if c.args and "refund paid" in c.args[0]
] ]
assert len(token_issued_calls) == 1 assert len(paid_calls) == 1
extra = token_issued_calls[0].kwargs.get("extra", {}) extra = paid_calls[0].kwargs.get("extra", {})
assert extra.get("path") == "/v1/wallet/refund" assert extra.get("method") == "cashu"
assert extra.get("refund_id")
@pytest.mark.asyncio @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 # Debit returns rowcount=0 → balance changed concurrently
session.exec = AsyncMock(return_value=_update_result(0)) session.exec = AsyncMock(return_value=_update_result(0))
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted") mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted")
with ( with (
patch("routstr.balance.send_token", mock_send_token), patch("routstr.refund.send_token", mock_send_token),
patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.refund.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
): ):
with pytest.raises(HTTPException) as exc_info: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( await refund_wallet_endpoint(
@@ -427,6 +460,7 @@ async def test_credit_balance_stores_apikey_transaction_history() -> None:
session = MagicMock() session = MagicMock()
session.exec = AsyncMock(return_value=_update_result(1)) session.exec = AsyncMock(return_value=_update_result(1))
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
session.refresh = AsyncMock() session.refresh = AsyncMock()
with ( with (
@@ -460,18 +494,20 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
# First exec call = debit (succeeds), second = restore # First exec call = debit (succeeds), second = restore
session = MagicMock() session = MagicMock()
session.get = AsyncMock(return_value=key) 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.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch( patch(
"routstr.balance.send_token", "routstr.refund.send_token",
AsyncMock(side_effect=MintConnectionError("raw mint outage detail")), AsyncMock(side_effect=MintConnectionError("raw mint outage detail")),
), ),
patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.refund.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), patch("routstr.refund.logger"),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger"),
): ):
with pytest.raises(HTTPException) as exc_info: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( 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.status_code == 503
assert exc_info.value.detail == "Mint service unavailable" assert exc_info.value.detail == "Mint service unavailable"
assert "raw mint outage detail" not in exc_info.value.detail assert "raw mint outage detail" not in exc_info.value.detail
# Verify two exec calls: debit + restore # debit, claim close, balance restore
assert session.exec.await_count == 2 assert session.exec.await_count == 3
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -497,15 +533,19 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None:
session = MagicMock() session = MagicMock()
session.get = AsyncMock(return_value=key) 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.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))), patch(
patch("routstr.balance.store_cashu_transaction", AsyncMock()), "routstr.refund.send_token", AsyncMock(side_effect=RuntimeError(raw_error))
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)), ),
patch("routstr.balance._refund_cache_set", AsyncMock()), patch("routstr.refund.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.logger"), patch("routstr.refund.logger"),
): ):
with pytest.raises(HTTPException) as exc_info: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( 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.status_code == 500
assert exc_info.value.detail == "Refund failed" assert exc_info.value.detail == "Refund failed"
assert raw_error not in exc_info.value.detail 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) ------------------ # --- Topup redemption error taxonomy (POST /v1/wallet/topup) ------------------
def _envelope(exc: HTTPException) -> dict: def _envelope(exc: HTTPException) -> dict:
"""Extract the error object from a top-up HTTPException.""" """Extract the error object from a top-up HTTPException."""
detail = exc.detail detail = exc.detail
@@ -633,7 +674,9 @@ async def test_topup_mint_unreachable_returns_503(
@pytest.mark.asyncio @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 fastapi import HTTPException
from routstr.wallet import SourceMintConnectionError from routstr.wallet import SourceMintConnectionError
@@ -697,7 +740,9 @@ async def test_topup_zero_value_returns_400_zero_value_message() -> None:
patch( patch(
"routstr.balance.credit_balance", "routstr.balance.credit_balance",
AsyncMock( 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", "Token value is too small to cover swap fees",
), ),
( (
ValueError( ValueError("Token amount (5 sat) is insufficient to cover melt fees."),
"Token amount (5 sat) is insufficient to cover melt fees."
),
422, 422,
"mint_error", "mint_error",
"cashu_token_swap_fees_exceed_amount", "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.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) session.exec = AsyncMock(return_value=MagicMock(rowcount=1))
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch( patch(
"routstr.balance.send_to_lnurl", "routstr.refund.send_to_lnurl",
AsyncMock(side_effect=MeltOutcomeAmbiguousError("outcome is ambiguous")), 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: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( 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.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=MagicMock(rowcount=1)) session.exec = AsyncMock(return_value=MagicMock(rowcount=1))
session.commit = AsyncMock() session.commit = AsyncMock()
session.rollback = AsyncMock()
with ( with (
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch( patch(
"routstr.balance.send_to_lnurl", "routstr.refund.send_to_lnurl",
AsyncMock(side_effect=RuntimeError("mint rejected melt")), 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): with pytest.raises(HTTPException):
await refund_wallet_endpoint( await refund_wallet_endpoint(
+10 -12
View File
@@ -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] def _refund_patches(refund_token: str = "cashuArefund"): # type: ignore[no-untyped-def]
return ( return (
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)), patch("routstr.refund.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()), patch("routstr.refund.store_cashu_transaction", AsyncMock()),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
) )
@@ -224,8 +222,8 @@ async def test_refund_self_heals_stale_reservation(session: AsyncSession) -> Non
reserved_at=int(time.time()) - 10_000, reserved_at=int(time.time()) - 10_000,
) )
p1, p2, p3, p4 = _refund_patches() p1, p2 = _refund_patches()
with p1, p2, p3, p4: with p1, p2:
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-stalerefund", authorization="Bearer sk-stalerefund",
x_cashu=None, x_cashu=None,
@@ -253,8 +251,8 @@ async def test_refund_self_heals_legacy_null_reserved_at(session: AsyncSession)
reserved_at=None, reserved_at=None,
) )
p1, p2, p3, p4 = _refund_patches() p1, p2 = _refund_patches()
with p1, p2, p3, p4: with p1, p2:
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-legacyrefund", authorization="Bearer sk-legacyrefund",
x_cashu=None, x_cashu=None,
@@ -280,8 +278,8 @@ async def test_refund_rejects_recent_reservation(session: AsyncSession) -> None:
reserved_at=int(time.time()), reserved_at=int(time.time()),
) )
p1, p2, p3, p4 = _refund_patches() p1, p2 = _refund_patches()
with p1, p2, p3, p4: with p1, p2:
with pytest.raises(HTTPException) as exc_info: with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint( await refund_wallet_endpoint(
authorization="Bearer sk-activerefund", authorization="Bearer sk-activerefund",
@@ -302,8 +300,8 @@ async def test_refund_without_reservation_still_works(session: AsyncSession) ->
reserved_balance=0, reserved_balance=0,
) )
p1, p2, p3, p4 = _refund_patches() p1, p2 = _refund_patches()
with p1, p2, p3, p4: with p1, p2:
result = await refund_wallet_endpoint( result = await refund_wallet_endpoint(
authorization="Bearer sk-plainrefund", authorization="Bearer sk-plainrefund",
x_cashu=None, x_cashu=None,