This commit is contained in:
9qeklajc
2026-08-03 23:32:06 +02:00
parent dd8c4a9a8a
commit da859f2f84
25 changed files with 1342 additions and 239 deletions
@@ -5,6 +5,9 @@ Revises: 64ed5594df1f
Create Date: 2026-08-02 23:53:00.037456 Create Date: 2026-08-02 23:53:00.037456
""" """
import json
import os
import sqlalchemy as sa import sqlalchemy as sa
from alembic import op from alembic import op
@@ -15,11 +18,50 @@ branch_labels = None
depends_on = None depends_on = None
def _resolve_backfill_mint_url(bind: sa.engine.Connection) -> str | None:
"""Best-effort resolution of the mint that issued pre-existing invoices.
Order: persisted settings JSON -> PRIMARY_MINT_URL env -> first CASHU_MINTS entry.
"""
try:
row = bind.execute(
sa.text("SELECT data FROM settings ORDER BY id LIMIT 1")
).fetchone()
if row and row[0]:
data = json.loads(row[0])
mint = data.get("primary_mint") or next(
iter(data.get("cashu_mints") or []), None
)
if mint:
return str(mint)
except Exception:
pass
env_mint = os.environ.get("PRIMARY_MINT_URL", "").strip()
if env_mint:
return env_mint
cashu_mints = os.environ.get("CASHU_MINTS", "").strip()
if cashu_mints:
return cashu_mints.split(",")[0].strip() or None
return None
def upgrade() -> None: def upgrade() -> None:
op.add_column( op.add_column(
"lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True) "lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True)
) )
bind = op.get_bind()
backfill_mint = _resolve_backfill_mint_url(bind)
if backfill_mint:
bind.execute(
sa.text(
"UPDATE lightning_invoices SET mint_url = :mint WHERE mint_url IS NULL"
),
{"mint": backfill_mint},
)
def downgrade() -> None: def downgrade() -> None:
op.drop_column("lightning_invoices", "mint_url") op.drop_column("lightning_invoices", "mint_url")
+14 -2
View File
@@ -387,11 +387,23 @@ async def _validate_bearer_key_locked(
"has_expiry_time": bool(key_expiry_time), "has_expiry_time": bool(key_expiry_time),
}, },
) )
if token_obj.mint in settings.cashu_mints: if token_obj.mint == settings.primary_mint:
if token_obj.unit != settings.primary_mint_unit:
raise redemption_error_to_http_exception(
ValueError(
"Cashu token unit does not match the configured primary "
f"mint unit: expected {settings.primary_mint_unit}, "
f"got {token_obj.unit}"
)
)
refund_currency = token_obj.unit
refund_mint_url = settings.primary_mint
elif token_obj.mint in settings.cashu_mints:
refund_currency = token_obj.unit refund_currency = token_obj.unit
refund_mint_url = token_obj.mint refund_mint_url = token_obj.mint
else: else:
refund_currency = "sat" # Foreign tokens are swapped into the configured primary mint.
refund_currency = settings.primary_mint_unit
refund_mint_url = settings.primary_mint refund_mint_url = settings.primary_mint
new_key = ApiKey( new_key = ApiKey(
+15 -27
View File
@@ -13,13 +13,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession
from ..payment.models import _row_to_model, list_models from ..payment.models import _row_to_model, list_models
from ..proxy import refresh_model_maps, reinitialize_upstreams from ..proxy import refresh_model_maps, reinitialize_upstreams
from ..wallet import ( from ..wallet import fetch_all_balances, send_token, token_mint_url
fetch_all_balances,
get_proofs_per_mint_and_unit,
get_wallet,
send_token,
slow_filter_spend_proofs,
)
from . import vault from . import vault
from .db import ( from .db import (
ApiKey, ApiKey,
@@ -442,37 +436,31 @@ class WithdrawRequest(BaseModel):
async def withdraw( async def withdraw(
request: Request, withdraw_request: WithdrawRequest request: Request, withdraw_request: WithdrawRequest
) -> dict[str, str]: ) -> dict[str, str]:
# Get wallet and check balance
from .settings import settings as global_settings from .settings import settings as global_settings
effective_mint = withdraw_request.mint_url or global_settings.primary_mint effective_mint = withdraw_request.mint_url or global_settings.primary_mint
wallet = await get_wallet(effective_mint, withdraw_request.unit)
proofs = get_proofs_per_mint_and_unit(
wallet,
effective_mint,
withdraw_request.unit,
not_reserved=True,
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
current_balance = sum(proof.amount for proof in proofs)
if withdraw_request.amount <= 0: if withdraw_request.amount <= 0:
raise HTTPException( raise HTTPException(
status_code=400, detail="Withdrawal amount must be positive" status_code=400, detail="Withdrawal amount must be positive"
) )
if withdraw_request.amount > current_balance: try:
raise HTTPException(status_code=400, detail="Insufficient wallet balance") token = await send_token(
withdraw_request.amount, withdraw_request.unit, effective_mint
token = await send_token( )
withdraw_request.amount, withdraw_request.unit, effective_mint except ValueError as error:
) if not str(error).startswith("No trusted mint has "):
raise
raise HTTPException(
status_code=400, detail="Insufficient wallet balance"
) from error
actual_mint = token_mint_url(token, effective_mint)
try: try:
await store_cashu_transaction( await store_cashu_transaction(
token=token, token=token,
amount=withdraw_request.amount, amount=withdraw_request.amount,
unit=withdraw_request.unit, unit=withdraw_request.unit,
mint_url=effective_mint, mint_url=actual_mint,
typ="out", typ="out",
collected=False, collected=False,
source="admin", source="admin",
@@ -483,10 +471,10 @@ async def withdraw(
extra={ extra={
"amount": withdraw_request.amount, "amount": withdraw_request.amount,
"unit": withdraw_request.unit, "unit": withdraw_request.unit,
"mint_url": effective_mint, "mint_url": actual_mint,
}, },
) )
return {"token": token} return {"token": token, "mint_url": actual_mint}
class ModelCreate(BaseModel): class ModelCreate(BaseModel):
+8 -3
View File
@@ -293,7 +293,7 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
"""Delete dead parentless API keys; return the count removed. """Delete dead parentless API keys; return the count removed.
Dead = 0 balance/reservation/spend/requests, older than the grace period, Dead = 0 balance/reservation/spend/requests, older than the grace period,
no parent, no children, no pending invoice. Cashu rows are unlinked (not no parent, no children, no retryable invoice. Cashu rows are unlinked (not
deleted) first to keep the audit trail. deleted) first to keep the audit trail.
""" """
cutoff = int(time.time()) - min_age_seconds cutoff = int(time.time()) - min_age_seconds
@@ -307,7 +307,9 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
pending_invoice = ( pending_invoice = (
select(LightningInvoice.id) select(LightningInvoice.id)
.where(col(LightningInvoice.api_key_hash) == col(ApiKey.hashed_key)) .where(col(LightningInvoice.api_key_hash) == col(ApiKey.hashed_key))
.where(col(LightningInvoice.status) == "pending") .where(
col(LightningInvoice.status).in_(("pending", "settlement_pending"))
)
).exists() ).exists()
eligible_hashes = ( eligible_hashes = (
@@ -435,7 +437,10 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
payment_hash: str = Field(description="Payment hash for tracking", unique=True) payment_hash: str = Field(description="Payment hash for tracking", unique=True)
status: str = Field( status: str = Field(
default="pending", default="pending",
description="pending, paid, expired, cancelled, reconciliation_required", description=(
"pending, settlement_pending, paid, expired, cancelled, "
"reconciliation_required"
),
) )
api_key_hash: str | None = Field( api_key_hash: str | None = Field(
default=None, description="Associated API key hash for topup operations" default=None, description="Associated API key hash for topup operations"
+163 -42
View File
@@ -7,6 +7,7 @@ from contextlib import asynccontextmanager
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, AsyncGenerator from typing import Any, AsyncGenerator
from cashu.core.base import MintQuoteState
from fastapi import APIRouter, Depends, Header, HTTPException from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sqlalchemy.orm.attributes import set_committed_value from sqlalchemy.orm.attributes import set_committed_value
@@ -32,8 +33,8 @@ logger = get_logger(__name__)
lightning_router = APIRouter(prefix="/lightning") lightning_router = APIRouter(prefix="/lightning")
# Avoid duplicate work within one process. Cross-process credit fencing is done # Avoid duplicate work within one process. Cross-process settlement is fenced
# by the conditional pending -> paid update in _finalize_invoice_settlement(). # by claiming a paid quote before minting and by the final conditional update.
@dataclass @dataclass
class _InvoiceLockEntry: class _InvoiceLockEntry:
lock: asyncio.Lock lock: asyncio.Lock
@@ -138,6 +139,9 @@ class InvoiceStatusResponse(BaseModel):
expires_at: int expires_at: int
_RETRYABLE_INVOICE_STATUSES = ("pending", "settlement_pending")
class InvoiceRecoverRequest(BaseModel): class InvoiceRecoverRequest(BaseModel):
bolt11: str = Field(description="BOLT11 invoice string") bolt11: str = Field(description="BOLT11 invoice string")
@@ -166,8 +170,11 @@ async def _request_mint_with_fallback(
f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}." f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}."
) )
tried: list[str] = [] tried: list[str] = []
configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints] candidates = (
candidates = list(dict.fromkeys(configured)) list(dict.fromkeys(allowed_mints))
if allowed_mints
else _trusted_mint_candidates()
)
for mint_url in candidates: for mint_url in candidates:
cooldown = mint_cooldown_remaining(mint_url) cooldown = mint_cooldown_remaining(mint_url)
if cooldown > 0: if cooldown > 0:
@@ -317,12 +324,12 @@ async def get_invoice_status(
if not invoice: if not invoice:
raise HTTPException(status_code=404, detail="Invoice not found") raise HTTPException(status_code=404, detail="Invoice not found")
if invoice.status == "pending": definitively_unpaid = False
await check_invoice_payment(invoice, session) if invoice.status in _RETRYABLE_INVOICE_STATUSES:
definitively_unpaid = await check_invoice_payment(invoice, session)
if invoice.status == "pending" and int(time.time()) > invoice.expires_at: await _expire_invoice_if_authoritatively_unpaid(
invoice.status = "expired" invoice, session, definitively_unpaid
await session.commit() )
api_key = None api_key = None
if invoice.status == "paid" and invoice.purpose == "create": if invoice.status == "paid" and invoice.purpose == "create":
@@ -356,8 +363,12 @@ async def recover_invoice(
if not invoice: if not invoice:
raise HTTPException(status_code=404, detail="Invoice not found") raise HTTPException(status_code=404, detail="Invoice not found")
if invoice.status == "pending": definitively_unpaid = False
await check_invoice_payment(invoice, session) if invoice.status in _RETRYABLE_INVOICE_STATUSES:
definitively_unpaid = await check_invoice_payment(invoice, session)
await _expire_invoice_if_authoritatively_unpaid(
invoice, session, definitively_unpaid
)
api_key = None api_key = None
if invoice.status == "paid": if invoice.status == "paid":
@@ -376,19 +387,58 @@ async def recover_invoice(
) )
async def _claim_paid_invoice_for_settlement(
invoice: LightningInvoice,
caller_session: AsyncSession,
observed_status: str,
) -> bool:
"""Claim an authoritative paid quote before consuming it at the mint."""
if observed_status == "settlement_pending":
return True
if observed_status != "pending":
await _reload_invoice_view(invoice, caller_session)
return False
async with create_session() as claim_session:
claim = await claim_session.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(
col(LightningInvoice.id) == invoice.id,
col(LightningInvoice.status) == "pending",
)
.values(status="settlement_pending")
.execution_options(synchronize_session=False)
)
await claim_session.commit()
if claim.rowcount != 1:
await _reload_invoice_view(invoice, caller_session)
return False
_publish_invoice_value(invoice, "status", "settlement_pending")
return True
async def check_invoice_payment( async def check_invoice_payment(
invoice: LightningInvoice, session: AsyncSession invoice: LightningInvoice, session: AsyncSession
) -> None: ) -> bool:
"""Settle an invoice and report whether its quote is definitively unpaid.
False covers paid, pending, and ambiguous transport/DB outcomes so callers
never expire a quote merely because reconciliation could not complete.
"""
async with _invoice_settlement_lock(invoice.id), wallet_operation_guard(): async with _invoice_settlement_lock(invoice.id), wallet_operation_guard():
minted = False minted = False
payment_confirmed = False
try: try:
# Snapshot the row and end the caller's read transaction before any # Snapshot the row and end the caller's read transaction before any
# potentially slow mint I/O. All final DB mutations use owned, # potentially slow mint I/O. All final DB mutations use owned,
# short-lived sessions below. # short-lived sessions below.
await session.refresh(invoice) await session.refresh(invoice)
if invoice.status != "pending": if invoice.status not in _RETRYABLE_INVOICE_STATUSES:
await session.commit() await session.commit()
return return False
observed_status = invoice.status
settlement = _InvoiceSettlement.from_invoice(invoice) settlement = _InvoiceSettlement.from_invoice(invoice)
await session.commit() await session.commit()
@@ -400,7 +450,15 @@ async def check_invoice_payment(
mint_url=mint_url, mint_url=mint_url,
) )
if not mint_status.paid: if not mint_status.paid:
return return getattr(mint_status, "state", None) == MintQuoteState.unpaid
payment_confirmed = True
# Fence expiry and other workers before consuming the paid quote.
# If a concurrent expiry/finalization won, this worker must not mint.
if not await _claim_paid_invoice_for_settlement(
invoice, session, observed_status
):
return False
# Reject a paid top-up whose target was pruned before redeeming its # Reject a paid top-up whose target was pruned before redeeming its
# single-use quote. The validation session is closed before mint I/O. # single-use quote. The validation session is closed before mint I/O.
@@ -416,7 +474,9 @@ async def check_invoice_payment(
update(LightningInvoice) update(LightningInvoice)
.where( .where(
col(LightningInvoice.id) == settlement.id, col(LightningInvoice.id) == settlement.id,
col(LightningInvoice.status) == "pending", col(LightningInvoice.status).in_(
_RETRYABLE_INVOICE_STATUSES
),
) )
.values(status="reconciliation_required") .values(status="reconciliation_required")
) )
@@ -431,7 +491,7 @@ async def check_invoice_payment(
"Paid topup invoice target API key was not found; reconciliation required", "Paid topup invoice target API key was not found; reconciliation required",
extra={"invoice_id": settlement.id}, extra={"invoice_id": settlement.id},
) )
return return False
# Quote-linked proof verification makes an ambiguous mint response # Quote-linked proof verification makes an ambiguous mint response
# retryable without crediting unrelated wallet balance growth. # retryable without crediting unrelated wallet balance growth.
@@ -445,7 +505,7 @@ async def check_invoice_payment(
) )
if not settled: if not settled:
await _reload_invoice_view(invoice, session) await _reload_invoice_view(invoice, session)
return return False
_publish_invoice_value(invoice, "status", "paid") _publish_invoice_value(invoice, "status", "paid")
_publish_invoice_value(invoice, "paid_at", paid_at) _publish_invoice_value(invoice, "paid_at", paid_at)
@@ -461,9 +521,33 @@ async def check_invoice_payment(
else None, else None,
}, },
) )
return False
except BaseException as error: except BaseException as error:
# Never roll back the caller-owned session: doing so expires invoice # Never roll back the caller-owned session: doing so expires invoice
# and sibling ORM objects. Owned sessions roll themselves back. # and sibling ORM objects. Owned sessions roll themselves back.
if payment_confirmed and invoice.status != "settlement_pending":
try:
async with create_session() as state_session:
pending = await state_session.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(
col(LightningInvoice.id) == invoice.id,
col(LightningInvoice.status).in_(
_RETRYABLE_INVOICE_STATUSES
),
)
.values(status="settlement_pending")
)
await state_session.commit()
if pending.rowcount == 1:
_publish_invoice_value(
invoice, "status", "settlement_pending"
)
except Exception as state_error:
logger.critical(
"Paid invoice reconciliation state could not be persisted",
extra={"invoice_id": invoice.id, "error": str(state_error)},
)
if minted: if minted:
logger.critical( logger.critical(
"Invoice mint succeeded but DB finalization failed; reconciliation required", "Invoice mint succeeded but DB finalization failed; reconciliation required",
@@ -476,6 +560,7 @@ async def check_invoice_payment(
if not isinstance(error, Exception): if not isinstance(error, Exception):
raise raise
logger.error(f"Failed to check invoice payment: {error}") logger.error(f"Failed to check invoice payment: {error}")
return False
def _is_outputs_already_signed(error: BaseException) -> bool: def _is_outputs_already_signed(error: BaseException) -> bool:
@@ -588,7 +673,9 @@ async def _finalize_invoice_settlement(
claim = await session.exec( # type: ignore[call-overload] claim = await session.exec( # type: ignore[call-overload]
update(LightningInvoice) update(LightningInvoice)
.where(col(LightningInvoice.id) == invoice.id) .where(col(LightningInvoice.id) == invoice.id)
.where(col(LightningInvoice.status) == "pending") .where(
col(LightningInvoice.status).in_(_RETRYABLE_INVOICE_STATUSES)
)
.values(status="paid", paid_at=paid_at, api_key_hash=api_key_hash) .values(status="paid", paid_at=paid_at, api_key_hash=api_key_hash)
.execution_options(synchronize_session=False) .execution_options(synchronize_session=False)
) )
@@ -623,6 +710,39 @@ async def _reload_invoice_view(
_publish_invoice_value(invoice, "api_key_hash", api_key_hash) _publish_invoice_value(invoice, "api_key_hash", api_key_hash)
async def _expire_invoice_if_authoritatively_unpaid(
invoice: LightningInvoice,
caller_session: AsyncSession,
definitively_unpaid: bool,
) -> bool:
"""Expire one overdue unpaid invoice without overwriting concurrent settlement."""
if (
not definitively_unpaid
or invoice.status != "pending"
or int(time.time()) <= invoice.expires_at
):
return False
async with create_session() as expiry_session:
expired = await expiry_session.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(
col(LightningInvoice.id) == invoice.id,
col(LightningInvoice.status) == "pending",
)
.values(status="expired")
.execution_options(synchronize_session=False)
)
await expiry_session.commit()
if expired.rowcount == 1:
_publish_invoice_value(invoice, "status", "expired")
return True
await _reload_invoice_view(invoice, caller_session)
return False
async def _credit_topup_record( async def _credit_topup_record(
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
) -> None: ) -> None:
@@ -635,32 +755,33 @@ INVOICE_WATCH_INTERVAL_SECONDS = 10
INVOICE_WATCH_BATCH_LIMIT = 100 INVOICE_WATCH_BATCH_LIMIT = 100
async def periodic_invoice_watcher() -> None: async def _process_invoice_watch_batch(session: AsyncSession) -> None:
"""Background task: detect paid Lightning invoices and credit balances. result = await session.exec(
select(LightningInvoice)
.where(
col(LightningInvoice.status).in_(_RETRYABLE_INVOICE_STATUSES)
)
.limit(INVOICE_WATCH_BATCH_LIMIT)
)
for invoice in result.all():
try:
definitively_unpaid = await check_invoice_payment(invoice, session)
await _expire_invoice_if_authoritatively_unpaid(
invoice, session, definitively_unpaid
)
except Exception as e:
logger.error(
"Invoice watcher failed for invoice",
extra={"invoice_id": invoice.id, "error": str(e)},
)
Removes the need for clients to poll the status endpoint after paying.
""" async def periodic_invoice_watcher() -> None:
"""Background task: detect paid Lightning invoices and credit balances."""
while True: while True:
try: try:
async with create_session() as session: async with create_session() as session:
now = int(time.time()) await _process_invoice_watch_batch(session)
result = await session.exec(
select(LightningInvoice)
.where(
LightningInvoice.status == "pending",
col(LightningInvoice.expires_at) > now,
)
.limit(INVOICE_WATCH_BATCH_LIMIT)
)
pending = result.all()
for invoice in pending:
try:
await check_invoice_payment(invoice, session)
except Exception as e:
logger.error(
"Invoice watcher failed for invoice",
extra={"invoice_id": invoice.id, "error": str(e)},
)
except asyncio.CancelledError: except asyncio.CancelledError:
raise raise
except Exception as e: except Exception as e:
+14 -9
View File
@@ -72,7 +72,15 @@ class MintRateGuard:
concurrency = settings.mint_max_concurrency concurrency = settings.mint_max_concurrency
guard = cls._guards.get(mint_url) guard = cls._guards.get(mint_url)
if guard is None or guard._max_concurrency != concurrency: if guard is None or guard._max_concurrency != concurrency:
previous = guard
guard = cls(mint_url, concurrency) guard = cls(mint_url, concurrency)
if previous is not None:
# Concurrency changed at runtime: keep the live cooldown/backoff
# state so an active 429 cooldown is not silently discarded.
guard._cooldown_until = previous._cooldown_until
guard._cooldown_reason = previous._cooldown_reason
guard._consecutive_rate_limits = previous._consecutive_rate_limits
guard._needs_probe = previous._needs_probe
cls._guards[mint_url] = guard cls._guards[mint_url] = guard
return guard return guard
@@ -124,10 +132,9 @@ class MintRateGuard:
return self._cooldown_reason if self.cooldown_remaining() > 0 else None return self._cooldown_reason if self.cooldown_remaining() > 0 else None
def _raise_if_wait_forbidden(self) -> None: def _raise_if_wait_forbidden(self) -> None:
if _fail_fast_depth.get() and ( remaining = self.cooldown_remaining()
self._needs_probe or self.cooldown_remaining() > 0 if _fail_fast_depth.get() and remaining > 0:
): raise MintCooldownError(self._mint_url, remaining)
raise MintCooldownError(self._mint_url, self.cooldown_remaining())
async def _wait_for_cooldown(self) -> None: async def _wait_for_cooldown(self) -> None:
while True: while True:
@@ -157,11 +164,7 @@ class MintRateGuard:
retry_after = None retry_after = None
if isinstance(error, httpx.HTTPStatusError): if isinstance(error, httpx.HTTPStatusError):
retry_after = parse_retry_after(error.response.headers) retry_after = parse_retry_after(error.response.headers)
delay = max( self.apply_rate_limit_cooldown(retry_after)
_MINT_RATE_LIMIT_BASE_COOLDOWN_SECONDS,
retry_after or 0.0,
)
self.apply_cooldown(delay, reason="rate_limited")
else: else:
self.apply_cooldown(1.0) self.apply_cooldown(1.0)
logger.warning( logger.warning(
@@ -194,6 +197,8 @@ class MintRateGuard:
while True: while True:
self._raise_if_wait_forbidden() self._raise_if_wait_forbidden()
if self._needs_probe or self.cooldown_remaining() > 0: if self._needs_probe or self.cooldown_remaining() > 0:
if _fail_fast_depth.get() and self._probe_lock.locked():
raise MintCooldownError(self._mint_url, self.cooldown_remaining())
async with self._probe_lock: async with self._probe_lock:
self._raise_if_wait_forbidden() self._raise_if_wait_forbidden()
if self.cooldown_remaining() > 0: if self.cooldown_remaining() > 0:
+1 -1
View File
@@ -242,7 +242,7 @@ async def calculate_discounted_max_cost(
}, },
) )
return max(0, adjusted) return max(settings.min_request_msat, adjusted)
def estimate_tokens(messages: list) -> int: def estimate_tokens(messages: list) -> int:
+14 -2
View File
@@ -7,7 +7,11 @@ import httpx
from cashu.core.base import MeltQuoteState from cashu.core.base import MeltQuoteState
from cashu.wallet.wallet import Proof, Wallet from cashu.wallet.wallet import Proof, Wallet
from ..mint import MINT_TRANSPORT_EXCEPTIONS, run_mint_operation from ..mint import (
MINT_TRANSPORT_EXCEPTIONS,
is_mint_rate_limited,
run_mint_operation,
)
try: try:
from bech32 import bech32_decode, convertbits # type: ignore from bech32 import bech32_decode, convertbits # type: ignore
@@ -239,7 +243,15 @@ async def raw_send_to_lnurl(
mint_url=str(wallet.url), mint_url=str(wallet.url),
retry_timeouts=False, retry_timeouts=False,
) )
except MINT_TRANSPORT_EXCEPTIONS as error: except Exception as error:
if is_mint_rate_limited(error):
# Cooldown failures happen before dispatch, and HTTP 429 means the
# mint rejected the request. Neither outcome may keep proofs
# reserved as though a Lightning payment could still settle.
await wallet.set_reserved_for_send(proofs, reserved=False)
raise
if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS):
raise
melt_response = None melt_response = None
melt_error: BaseException | None = error melt_error: BaseException | None = error
else: else:
+230 -95
View File
@@ -310,18 +310,10 @@ async def _redeem_same_mint(
op_name="redeem_load_mint", op_name="redeem_load_mint",
mint_url=token_obj.mint, mint_url=token_obj.mint,
) )
wallet.verify_proofs_dleq(token_obj.proofs)
input_fees = wallet.get_fees_for_proofs(token_obj.proofs)
await run_mint_operation(
lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True),
op_name="redeem_split",
mint_url=token_obj.mint,
retry_timeouts=False,
)
except Exception as error: except Exception as error:
if is_mint_connection_error(error): if is_mint_connection_error(error):
logger.warning( logger.warning(
"Same-mint redemption failed; client must use a different token", "Same-mint redemption failed before swap dispatch",
extra={ extra={
"event": "cashu_same_mint_redemption_failed", "event": "cashu_same_mint_redemption_failed",
"source_mint": token_obj.mint, "source_mint": token_obj.mint,
@@ -338,23 +330,79 @@ async def _redeem_same_mint(
) from error ) from error
raise raise
wallet.verify_proofs_dleq(token_obj.proofs)
input_fees = wallet.get_fees_for_proofs(token_obj.proofs)
try:
await run_mint_operation(
lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True),
op_name="redeem_split",
mint_url=token_obj.mint,
retry_timeouts=False,
)
except Exception as error:
if isinstance(error, httpx.ConnectError):
raise SourceMintConnectionError(
"Issuing Cashu mint is unreachable"
) from error
if is_mint_connection_error(error):
logger.critical(
"Same-mint swap outcome is ambiguous; sealing source token",
extra={
"event": "cashu_same_mint_redemption_ambiguous",
"source_mint": token_obj.mint,
"source_unit": token_obj.unit,
"source_amount": token_obj.amount,
"action": "manual_reconciliation_required",
"error": str(error),
"error_type": type(error).__name__,
},
)
raise TokenConsumedError(
"Same-mint swap outcome is ambiguous; reconciliation required"
) from error
raise
return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint
async def recieve_token( async def recieve_token(
token: str, token: str,
destination_mint: str | None = None,
destination_unit: str | None = None,
) -> tuple[int, str, str]: # amount, unit, mint_url ) -> tuple[int, str, str]: # amount, unit, mint_url
"""Redeem a token while serializing all wallet proof mutation."""
async with wallet_operation_guard():
return await _recieve_token_locked(token, destination_mint, destination_unit)
async def _recieve_token_locked(
token: str,
destination_mint: str | None = None,
destination_unit: str | None = None,
) -> tuple[int, str, str]:
token_obj = deserialize_token_from_string(token) token_obj = deserialize_token_from_string(token)
if len(token_obj.keysets) > 1: if len(token_obj.keysets) > 1:
raise ValueError("Multiple keysets per token currently not supported") raise ValueError("Multiple keysets per token currently not supported")
destinations = (
[destination_mint]
if destination_mint is not None
else list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints]))
)
output_unit = (
token_obj.unit
if token_obj.mint in destinations
else settings.primary_mint_unit
)
if destination_unit is not None and output_unit != destination_unit:
raise ValueError(
"Cashu token unit does not match the API key liability unit: "
f"expected {destination_unit}, got {output_unit}"
)
wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False)
wallet.keyset_id = token_obj.keysets[0] wallet.keyset_id = token_obj.keysets[0]
if token_obj.mint not in destinations:
if token_obj.mint not in settings.cashu_mints:
destinations = list(
dict.fromkeys([settings.primary_mint, *settings.cashu_mints])
)
logger.info( logger.info(
"Cashu cross-mint swap required", "Cashu cross-mint swap required",
extra={ extra={
@@ -365,7 +413,9 @@ async def recieve_token(
"destination_candidates": destinations, "destination_candidates": destinations,
}, },
) )
return await swap_to_trusted_mint(token_obj, wallet) return await swap_to_trusted_mint(
token_obj, wallet, destination_mints=destinations
)
logger.info( logger.info(
"Trying same-mint Cashu redemption", "Trying same-mint Cashu redemption",
@@ -382,7 +432,16 @@ async def recieve_token(
async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]: async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]:
"""Create a token from the preferred mint or another funded trusted mint.""" """Create a token from the preferred mint or another funded trusted mint."""
effective_mint_url = await find_trusted_mint_with_funds(amount, unit, mint_url) async with wallet_operation_guard():
return await _send_locked(amount, unit, mint_url)
async def _send_locked(
amount: int, unit: str, mint_url: str | None = None
) -> tuple[int, str]:
effective_mint_url = await find_trusted_mint_with_funds(
amount, unit, mint_url, force_reload=True
)
wallet = await get_wallet(effective_mint_url, unit) wallet = await get_wallet(effective_mint_url, unit)
proofs = get_proofs_per_mint_and_unit( proofs = get_proofs_per_mint_and_unit(
wallet, effective_mint_url, unit, not_reserved=True wallet, effective_mint_url, unit, not_reserved=True
@@ -436,16 +495,20 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str
async def release_token_reservation(token: str) -> None: async def release_token_reservation(token: str) -> None:
"""Release a token that was created locally but never handed off.""" """Release a token that was created locally but never handed off."""
token_obj = deserialize_token_from_string(token) async with wallet_operation_guard():
wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False) token_obj = deserialize_token_from_string(token)
await wallet.set_reserved_for_send(token_obj.proofs, reserved=False) wallet = await get_wallet(token_obj.mint, token_obj.unit, load=False)
# This is a local wallet-DB refresh; reservation release must still work
# while the mint is unavailable or cooling down.
await wallet.load_proofs(reload=True)
await wallet.set_reserved_for_send(token_obj.proofs, reserved=False)
secrets = {proof.secret for proof in token_obj.proofs} secrets = {proof.secret for proof in token_obj.proofs}
for proof in token_obj.proofs: for proof in token_obj.proofs:
proof.reserved = False
for proof in wallet.proofs:
if proof.secret in secrets:
proof.reserved = False proof.reserved = False
for proof in wallet.proofs:
if proof.secret in secrets:
proof.reserved = False
def token_mint_url(token: str, fallback: str | None = None) -> str: def token_mint_url(token: str, fallback: str | None = None) -> str:
@@ -458,7 +521,11 @@ def token_mint_url(token: str, fallback: str | None = None) -> str:
async def find_trusted_mint_with_funds( async def find_trusted_mint_with_funds(
amount: int, unit: str, preferred_mint: str | None = None amount: int,
unit: str,
preferred_mint: str | None = None,
*,
force_reload: bool = False,
) -> str: ) -> str:
"""Choose a trusted mint that can cover a refund without waiting on cooldown.""" """Choose a trusted mint that can cover a refund without waiting on cooldown."""
trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints]))
@@ -472,7 +539,12 @@ async def find_trusted_mint_with_funds(
if mint_cooldown_remaining(mint_url) > 0: if mint_cooldown_remaining(mint_url) > 0:
continue continue
try: try:
wallet = await get_wallet(mint_url, unit, retry_on_rate_limit=False) wallet = await get_wallet(
mint_url,
unit,
retry_on_rate_limit=False,
force_reload=force_reload,
)
except Exception as error: except Exception as error:
if is_mint_connection_error(error) or is_mint_rate_limited(error): if is_mint_connection_error(error) or is_mint_rate_limited(error):
balances[mint_url] = 0 balances[mint_url] = 0
@@ -566,8 +638,27 @@ def _melt_insufficient_shortfall(error: Exception) -> int | None:
return 1 return 1
def _trusted_destination_candidates(
candidates: list[str] | None = None,
) -> list[str]:
trusted = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints]))
if candidates is None:
return trusted
selected = list(dict.fromkeys(candidates))
untrusted = [mint_url for mint_url in selected if mint_url not in trusted]
if untrusted:
raise ValueError(f"Untrusted destination mint: {untrusted[0]}")
if not selected:
raise ValueError("At least one trusted destination mint is required")
return selected
async def _request_mint_with_fallback( async def _request_mint_with_fallback(
amount: int, *, op_name: str, primary_wallet: Wallet | None = None amount: int,
*,
op_name: str,
primary_wallet: Wallet | None = None,
destination_mints: list[str] | None = None,
) -> tuple[Wallet, str, MintQuote]: ) -> tuple[Wallet, str, MintQuote]:
"""Try request_mint on the primary mint, fall back to other trusted mints """Try request_mint on the primary mint, fall back to other trusted mints
on transport or rate-limit failure. Returns the wallet, mint_url, and quote. on transport or rate-limit failure. Returns the wallet, mint_url, and quote.
@@ -581,7 +672,7 @@ async def _request_mint_with_fallback(
f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. " f"_request_mint_with_fallback({op_name}): amount must be > 0, got {amount}. "
f"Token value is too small after fee deduction or unit conversion." f"Token value is too small after fee deduction or unit conversion."
) )
candidates = list(dict.fromkeys([settings.primary_mint, *settings.cashu_mints])) candidates = _trusted_destination_candidates(destination_mints)
logger.info( logger.info(
"Trying trusted destination mints", "Trying trusted destination mints",
extra={ extra={
@@ -692,6 +783,7 @@ async def _calculate_swap_amount(
token_wallet: Wallet, token_wallet: Wallet,
primary_wallet: Wallet | None, primary_wallet: Wallet | None,
proofs: list, proofs: list,
destination_mints: list[str] | None = None,
) -> int: ) -> int:
""" """
Calculate the amount to mint on the primary mint after accounting for Calculate the amount to mint on the primary mint after accounting for
@@ -749,6 +841,7 @@ async def _calculate_swap_amount(
receive_amount, receive_amount,
op_name="swap_fee_est_mint_quote", op_name="swap_fee_est_mint_quote",
primary_wallet=primary_wallet, primary_wallet=primary_wallet,
destination_mints=destination_mints,
) )
stage = "source_fee_quote" stage = "source_fee_quote"
dummy_melt_quote = await run_mint_operation( dummy_melt_quote = await run_mint_operation(
@@ -869,7 +962,10 @@ async def _confirm_melt_paid(
async def swap_to_trusted_mint( async def swap_to_trusted_mint(
token_obj: Token, token_wallet: Wallet token_obj: Token,
token_wallet: Wallet,
*,
destination_mints: list[str] | None = None,
) -> tuple[int, str, str]: ) -> tuple[int, str, str]:
logger.info( logger.info(
"Starting Cashu cross-mint swap", "Starting Cashu cross-mint swap",
@@ -893,10 +989,11 @@ async def swap_to_trusted_mint(
amount_msat = token_amount amount_msat = token_amount
else: else:
raise ValueError("Invalid unit") raise ValueError("Invalid unit")
# If the token is already from the primary mint, we don't need a cross-mint destination_candidates = _trusted_destination_candidates(destination_mints)
# swap — redeem it same-mint. There's no melt/Lightning fee, but the mint's # If the token is already from an allowed destination, redeem it same-mint.
# NUT-02 input fee still applies; _redeem_same_mint accounts for it. # There's no melt/Lightning fee, but the mint's NUT-02 input fee still
if token_obj.mint == settings.primary_mint: # applies; _redeem_same_mint accounts for it.
if token_obj.mint in destination_candidates:
logger.info( logger.info(
"swap_to_trusted_mint: token already on primary mint, skipping swap", "swap_to_trusted_mint: token already on primary mint, skipping swap",
extra={ extra={
@@ -916,6 +1013,7 @@ async def swap_to_trusted_mint(
token_wallet, token_wallet,
primary_wallet, primary_wallet,
token_obj.proofs, token_obj.proofs,
destination_candidates,
) )
# The estimate above is non-binding: the mint may demand a higher fee on the # The estimate above is non-binding: the mint may demand a higher fee on the
@@ -949,6 +1047,7 @@ async def swap_to_trusted_mint(
minted_amount, minted_amount,
op_name="swap_request_mint", op_name="swap_request_mint",
primary_wallet=primary_wallet, primary_wallet=primary_wallet,
destination_mints=destination_candidates,
) )
logger.info( logger.info(
"swap_to_trusted_mint: mint quote received", "swap_to_trusted_mint: mint quote received",
@@ -1245,7 +1344,14 @@ async def _credit_balance_locked(
) )
try: try:
amount, unit, mint_url = await recieve_token(cashu_token) destination_mint = key.refund_mint_url or settings.primary_mint
amount, unit, mint_url = await recieve_token(
cashu_token,
destination_mint=destination_mint,
destination_unit=key.refund_currency
if isinstance(key.refund_currency, str)
else None,
)
original_amount = amount original_amount = amount
original_unit = unit original_unit = unit
logger.info( logger.info(
@@ -1284,10 +1390,19 @@ async def _credit_balance_locked(
# retryable/token-error taxonomy. # retryable/token-error taxonomy.
try: try:
# Atomic UPDATE to prevent race conditions during concurrent topups. # Atomic UPDATE to prevent race conditions during concurrent topups.
updates: dict[str, object] = {
"balance": db.ApiKey.balance + amount,
}
# Legacy keys may predate refund provenance. Pin them to the
# destination used for this credit before exposing the balance.
if key.refund_mint_url is None:
updates["refund_mint_url"] = mint_url
if key.refund_currency is None:
updates["refund_currency"] = unit
stmt = ( stmt = (
update(db.ApiKey) update(db.ApiKey)
.where(col(db.ApiKey.hashed_key) == key.hashed_key) .where(col(db.ApiKey.hashed_key) == key.hashed_key)
.values(balance=(db.ApiKey.balance) + amount) .values(**updates)
) )
result = await session.exec(stmt) # type: ignore[call-overload] result = await session.exec(stmt) # type: ignore[call-overload]
# If pruning removed this key after redemption, do not commit a no-op # If pruning removed this key after redemption, do not commit a no-op
@@ -1355,6 +1470,7 @@ async def get_wallet(
unit: str = "sat", unit: str = "sat",
load: bool = True, load: bool = True,
retry_on_rate_limit: bool = True, retry_on_rate_limit: bool = True,
force_reload: bool = False,
) -> Wallet: ) -> Wallet:
global _wallets, _wallet_last_load, _wallet_load_locks global _wallets, _wallet_last_load, _wallet_load_locks
id = f"{mint_url}_{unit}" id = f"{mint_url}_{unit}"
@@ -1366,7 +1482,11 @@ async def get_wallet(
if load: if load:
now = time.monotonic() now = time.monotonic()
last = _wallet_last_load.get(id) last = _wallet_last_load.get(id)
if last is None or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS: if (
force_reload
or last is None
or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS
):
await run_mint_operation( await run_mint_operation(
lambda: _wallets[id].load_mint(), lambda: _wallets[id].load_mint(),
op_name="load_mint", op_name="load_mint",
@@ -1849,7 +1969,8 @@ async def _refund_sweep_once(cutoff: int) -> None:
claim_owned = col(db.CashuTransaction.sweep_started_at) == claim_started_at claim_owned = col(db.CashuTransaction.sweep_started_at) == claim_started_at
redeemed = False redeemed = False
try: try:
await recieve_token(refund.token) async with wallet_operation_guard():
await recieve_token(refund.token)
redeemed = True redeemed = True
finalized = await _set_refund_sweep_state( finalized = await _set_refund_sweep_state(
refund.id, refund.id,
@@ -1980,66 +2101,73 @@ async def periodic_routstr_fee_payout() -> None:
continue continue
paid_msats = _sats_to_msats(accumulated_sats) paid_msats = _sats_to_msats(accumulated_sats)
# Wallet/proof preparation cannot send funds, so do it before the # Serialize proof refresh, reservation, sending, and checkpoint
# durable checkpoint. A preparation failure must not strand an # finalization with every other wallet mutation across workers.
# in-progress payout that requires manual reconciliation. async with wallet_operation_guard():
wallet = await get_wallet(settings.primary_mint, "sat") # Wallet/proof preparation cannot send funds, so do it before
proofs = get_proofs_per_mint_and_unit( # the durable checkpoint. Force a DB reload after taking the
wallet, settings.primary_mint, "sat", not_reserved=True # guard so another worker's reservations are visible.
) wallet = await get_wallet(
settings.primary_mint, "sat", force_reload=True
async with db.create_session() as session:
payout_checkpointed = await db.reset_routstr_fee(session, paid_msats)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
try:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
ROUTSTR_LN_ADDRESS,
"sat",
amount=accumulated_sats,
) )
except BaseException as e: proofs = get_proofs_per_mint_and_unit(
logger.critical( wallet, settings.primary_mint, "sat", not_reserved=True
"Routstr fee payout outcome is unknown; manual reconciliation required",
extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception),
) )
if not isinstance(e, Exception):
raise
continue
try:
async with db.create_session() as session: async with db.create_session() as session:
payout_completed = await db.complete_routstr_fee_payout( payout_checkpointed = await db.reset_routstr_fee(
session, paid_msats session, paid_msats
) )
except BaseException as e: if not payout_checkpointed:
logger.critical( logger.warning("Routstr fee payout was already claimed")
"Routstr fee payout sent but checkpoint was not completed", continue
extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
continue
if not payout_completed:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
)
continue
logger.info( try:
"Routstr fee payout sent", amount_received = await raw_send_to_lnurl(
extra={ wallet,
"accumulated_sats": accumulated_sats, proofs,
"amount_received": amount_received, ROUTSTR_LN_ADDRESS,
}, "sat",
) amount=accumulated_sats,
)
except BaseException as e:
logger.critical(
"Routstr fee payout outcome is unknown; manual reconciliation required",
extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
continue
try:
async with db.create_session() as session:
payout_completed = await db.complete_routstr_fee_payout(
session, paid_msats
)
except BaseException as e:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
continue
if not payout_completed:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
)
continue
logger.info(
"Routstr fee payout sent",
extra={
"accumulated_sats": accumulated_sats,
"amount_received": amount_received,
},
)
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Error in Routstr fee payout: {type(e).__name__}", f"Error in Routstr fee payout: {type(e).__name__}",
@@ -2048,11 +2176,18 @@ 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) -> int:
mint = await find_trusted_mint_with_funds(amount, unit, mint) async with wallet_operation_guard():
wallet = await get_wallet(mint, unit) mint = await find_trusted_mint_with_funds(
available = get_proofs_per_mint_and_unit(wallet, mint, unit, not_reserved=True) amount, unit, mint, force_reload=True
proofs, _ = await wallet.select_to_send(available, amount, set_reserved=True) )
return await raw_send_to_lnurl(wallet, proofs, address, unit) wallet = await get_wallet(mint, unit)
available = get_proofs_per_mint_and_unit(
wallet, mint, unit, not_reserved=True
)
proofs, _ = await wallet.select_to_send(
available, amount, set_reserved=True
)
return await raw_send_to_lnurl(wallet, proofs, address, unit)
# class Payment: # class Payment:
+7 -2
View File
@@ -203,8 +203,13 @@ class TestmintWallet:
token_base64 = base64.urlsafe_b64encode(token_json.encode()).decode() token_base64 = base64.urlsafe_b64encode(token_json.encode()).decode()
return f"cashuA{token_base64}" return f"cashuA{token_base64}"
async def redeem_token(self, token: str) -> Tuple[int, str, str]: async def redeem_token(
"""Redeem a Cashu token - compatible with wallet.recieve_token""" self,
token: str,
destination_mint: str | None = None,
destination_unit: str | None = None,
) -> Tuple[int, str, str]:
"""Redeem a Cashu token - compatible with wallet.recieve_token."""
if not self.wallet: if not self.wallet:
await self.init() await self.init()
@@ -246,7 +246,7 @@ async def test_concurrent_payment_checks_mint_and_credit_invoice_once(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_failed_mint_keeps_invoice_pending_for_retry( async def test_failed_mint_marks_invoice_for_settlement_retry(
integration_engine: AsyncEngine, integration_engine: AsyncEngine,
patched_db_engine: None, patched_db_engine: None,
) -> None: ) -> None:
@@ -270,7 +270,7 @@ async def test_failed_mint_keeps_invoice_pending_for_retry(
async with AsyncSession(integration_engine, expire_on_commit=False) as verify: async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
stored = await verify.get(LightningInvoice, invoice.id) stored = await verify.get(LightningInvoice, invoice.id)
assert stored is not None assert stored is not None
assert stored.status == "pending" assert stored.status == "settlement_pending"
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -399,14 +399,14 @@ async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation(
assert sibling_state is not None assert sibling_state is not None
assert stored_state.expired is False assert stored_state.expired is False
assert sibling_state.expired is False assert sibling_state.expired is False
assert stored.status == "pending" assert stored.status == "settlement_pending"
assert stored_sibling.id == sibling.id assert stored_sibling.id == sibling.id
assert wallet.mint.await_count == 1 assert wallet.mint.await_count == 1
async with AsyncSession(integration_engine, expire_on_commit=False) as verify: async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
stored = await verify.get(LightningInvoice, invoice.id) stored = await verify.get(LightningInvoice, invoice.id)
assert stored is not None assert stored is not None
assert stored.status == "pending" assert stored.status == "settlement_pending"
@pytest.mark.asyncio @pytest.mark.asyncio
+93 -1
View File
@@ -11,6 +11,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey, LightningInvoice from routstr.core.db import ApiKey, LightningInvoice
from routstr.lightning import ( from routstr.lightning import (
_expire_invoice_if_authoritatively_unpaid,
_finalize_invoice_settlement, _finalize_invoice_settlement,
_InvoiceSettlement, _InvoiceSettlement,
check_invoice_payment, check_invoice_payment,
@@ -251,7 +252,7 @@ async def test_check_invoice_payment_retries_after_mint_success_and_db_failure(
pending = await verify.get(LightningInvoice, invoice.id) pending = await verify.get(LightningInvoice, invoice.id)
unchanged = await verify.get(ApiKey, key_hash) unchanged = await verify.get(ApiKey, key_hash)
assert pending is not None assert pending is not None
assert pending.status == "pending" assert pending.status == "settlement_pending"
assert unchanged is not None assert unchanged is not None
assert unchanged.balance == 100_000 assert unchanged.balance == 100_000
@@ -271,3 +272,94 @@ async def test_check_invoice_payment_retries_after_mint_success_and_db_failure(
wallet.mint.assert_awaited_once_with(100, quote_id=invoice.payment_hash) wallet.mint.assert_awaited_once_with(100, quote_id=invoice.payment_hash)
wallet.restore_tokens_for_keyset.assert_not_awaited() wallet.restore_tokens_for_keyset.assert_not_awaited()
@pytest.mark.asyncio
async def test_expiry_cas_cannot_overwrite_concurrent_paid_invoice(
integration_engine: AsyncEngine,
patched_db_engine: None,
) -> None:
invoice = _lightning_invoice(expires_at=0)
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
seed.add(invoice)
await seed.commit()
async with AsyncSession(integration_engine, expire_on_commit=False) as caller:
stale = await caller.get(LightningInvoice, invoice.id)
assert stale is not None
await caller.commit()
async with AsyncSession(integration_engine, expire_on_commit=False) as paid:
result = await paid.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(col(LightningInvoice.id) == invoice.id)
.values(status="paid", paid_at=123)
)
assert result.rowcount == 1
await paid.commit()
expired = await _expire_invoice_if_authoritatively_unpaid(
stale, caller, True
)
assert expired is False
assert stale.status == "paid"
assert stale.paid_at == 123
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
stored = await verify.get(LightningInvoice, invoice.id)
assert stored is not None
assert stored.status == "paid"
assert stored.paid_at == 123
@pytest.mark.asyncio
async def test_paid_quote_worker_does_not_mint_after_expiry_claim_wins(
integration_engine: AsyncEngine,
patched_db_engine: None,
) -> None:
invoice = _lightning_invoice(expires_at=0)
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
seed.add(invoice)
await seed.commit()
quote_started = asyncio.Event()
release_quote = asyncio.Event()
async def paid_quote_after_expiry(*_args: object, **_kwargs: object) -> Mock:
quote_started.set()
await release_quote.wait()
return Mock(paid=True)
wallet = Mock(
get_mint_quote=AsyncMock(side_effect=paid_quote_after_expiry),
mint=AsyncMock(),
)
async with AsyncSession(integration_engine, expire_on_commit=False) as worker:
observed_pending = await worker.get(LightningInvoice, invoice.id)
assert observed_pending is not None
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
settlement_task = asyncio.create_task(
check_invoice_payment(observed_pending, worker)
)
await quote_started.wait()
async with AsyncSession(
integration_engine, expire_on_commit=False
) as expirer:
expiry_view = await expirer.get(LightningInvoice, invoice.id)
assert expiry_view is not None
await expirer.commit()
assert await _expire_invoice_if_authoritatively_unpaid(
expiry_view, expirer, True
)
release_quote.set()
assert await settlement_task is False
wallet.mint.assert_not_awaited()
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
stored = await verify.get(LightningInvoice, invoice.id)
assert stored is not None
assert stored.status == "expired"
@@ -110,7 +110,11 @@ async def test_payout_does_not_send_proofs_whose_liability_commit_is_in_flight(
finish_redemption = asyncio.Event() finish_redemption = asyncio.Event()
liability_read = asyncio.Event() liability_read = asyncio.Event()
async def redeem_token(token: str) -> tuple[int, str, str]: async def redeem_token(
token: str,
destination_mint: str | None = None,
destination_unit: str | None = None,
) -> tuple[int, str, str]:
proofs.append(MagicMock(amount=200)) proofs.append(MagicMock(amount=200))
proof_visible.set() proof_visible.set()
await finish_redemption.wait() await finish_redemption.wait()
@@ -126,8 +126,11 @@ async def test_parent_and_child_keys_are_not_pruned(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_pending_invoice_protects_key(patched_db_engine: None) -> None: @pytest.mark.parametrize("status", ["pending", "settlement_pending"])
"""A key referenced by a pending topup invoice is never pruned mid-topup.""" async def test_retryable_invoice_protects_key(
patched_db_engine: None, status: str
) -> None:
"""A key referenced by a retryable topup invoice is never pruned mid-topup."""
key = _dead_key(LONG_AGO) key = _dead_key(LONG_AGO)
invoice = LightningInvoice( invoice = LightningInvoice(
id=f"inv_{uuid.uuid4().hex}", id=f"inv_{uuid.uuid4().hex}",
@@ -135,7 +138,7 @@ async def test_pending_invoice_protects_key(patched_db_engine: None) -> None:
amount_sats=10, amount_sats=10,
description="topup", description="topup",
payment_hash=uuid.uuid4().hex, payment_hash=uuid.uuid4().hex,
status="pending", status=status,
api_key_hash=key.hashed_key, api_key_hash=key.hashed_key,
purpose="topup", purpose="topup",
expires_at=NOW + 10_000, expires_at=NOW + 10_000,
+3 -1
View File
@@ -29,7 +29,9 @@ from routstr.core.settings import settings
# with the testmint stub that bypasses swapping (see conftest.py). # with the testmint stub that bypasses swapping (see conftest.py).
from routstr.wallet import recieve_token as _real_recieve_token from routstr.wallet import recieve_token as _real_recieve_token
PRIMARY_MINT = "http://primary:3338" # Match the authenticated fixture's persisted refund mint: existing-key topups
# are intentionally constrained to that mint for collateral provenance.
PRIMARY_MINT = "http://localhost:3338"
def _make_swap_mocks( def _make_swap_mocks(
+92 -22
View File
@@ -1,8 +1,12 @@
import base64
import json
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock from unittest.mock import AsyncMock, Mock
import pytest import pytest
from fastapi import HTTPException
import routstr.wallet as wallet_module
from routstr.core import admin from routstr.core import admin
@@ -13,20 +17,12 @@ async def test_withdraw_uses_effective_mint_and_records_outgoing_transaction(
) -> None: ) -> None:
primary_mint = "https://primary.example" primary_mint = "https://primary.example"
effective_mint = requested_mint or primary_mint effective_mint = requested_mint or primary_mint
wallet = object()
proofs = [SimpleNamespace(amount=40), SimpleNamespace(amount=60)]
token = "cashuBoutgoing" token = "cashuBoutgoing"
get_wallet = AsyncMock(return_value=wallet)
get_proofs = Mock(return_value=proofs)
filter_proofs = AsyncMock(return_value=proofs)
send_token = AsyncMock(return_value=token) send_token = AsyncMock(return_value=token)
store_transaction = AsyncMock(return_value=True) store_transaction = AsyncMock(return_value=True)
monkeypatch.setattr(admin, "get_wallet", get_wallet)
monkeypatch.setattr(admin, "get_proofs_per_mint_and_unit", get_proofs)
monkeypatch.setattr(admin, "slow_filter_spend_proofs", filter_proofs)
monkeypatch.setattr(admin, "send_token", send_token) monkeypatch.setattr(admin, "send_token", send_token)
monkeypatch.setattr(admin, "token_mint_url", Mock(return_value=effective_mint))
monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction) monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction)
monkeypatch.setattr(admin.settings, "primary_mint", primary_mint) monkeypatch.setattr(admin.settings, "primary_mint", primary_mint)
@@ -35,10 +31,7 @@ async def test_withdraw_uses_effective_mint_and_records_outgoing_transaction(
admin.WithdrawRequest(amount=75, mint_url=requested_mint, unit="sat"), admin.WithdrawRequest(amount=75, mint_url=requested_mint, unit="sat"),
) )
assert result == {"token": token} assert result == {"token": token, "mint_url": effective_mint}
get_wallet.assert_awaited_once_with(effective_mint, "sat")
get_proofs.assert_called_once_with(wallet, effective_mint, "sat", not_reserved=True)
filter_proofs.assert_awaited_once_with(proofs, wallet)
send_token.assert_awaited_once_with(75, "sat", effective_mint) send_token.assert_awaited_once_with(75, "sat", effective_mint)
store_transaction.assert_awaited_once_with( store_transaction.assert_awaited_once_with(
token=token, token=token,
@@ -56,17 +49,10 @@ async def test_withdraw_returns_issued_token_when_audit_storage_fails(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
) -> None: ) -> None:
mint = "https://primary.example" mint = "https://primary.example"
proofs = [SimpleNamespace(amount=100)]
token = "cashuBrecoverable" token = "cashuBrecoverable"
monkeypatch.setattr(admin, "get_wallet", AsyncMock(return_value=object()))
monkeypatch.setattr(
admin, "get_proofs_per_mint_and_unit", Mock(return_value=proofs)
)
monkeypatch.setattr(
admin, "slow_filter_spend_proofs", AsyncMock(return_value=proofs)
)
monkeypatch.setattr(admin, "send_token", AsyncMock(return_value=token)) monkeypatch.setattr(admin, "send_token", AsyncMock(return_value=token))
monkeypatch.setattr(admin, "token_mint_url", Mock(return_value=mint))
monkeypatch.setattr( monkeypatch.setattr(
admin, admin,
"store_cashu_transaction", "store_cashu_transaction",
@@ -78,5 +64,89 @@ async def test_withdraw_returns_issued_token_when_audit_storage_fails(
result = await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75)) result = await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75))
assert result == {"token": token} assert result == {"token": token, "mint_url": mint}
critical.assert_called_once() critical.assert_called_once()
@pytest.mark.asyncio
async def test_withdraw_falls_back_from_insufficient_preferred_mint(
monkeypatch: pytest.MonkeyPatch,
) -> None:
requested_mint = "https://primary.example"
actual_mint = "https://secondary.example"
proofs = [SimpleNamespace(amount=100, reserved=False, id="00")]
token_payload = {
"token": [
{
"mint": actual_mint,
"proofs": [
{
"id": "00",
"amount": 75,
"secret": "secret",
"C": "02" + "00" * 32,
}
],
}
],
"unit": "sat",
}
token = "cashuA" + base64.urlsafe_b64encode(
json.dumps(token_payload).encode()
).decode()
wallet = SimpleNamespace(
keysets={},
proofs=proofs,
select_to_send=AsyncMock(return_value=(proofs, 0)),
serialize_proofs=AsyncMock(return_value=token),
set_reserved_for_send=AsyncMock(),
)
find_funded = AsyncMock(return_value=actual_mint)
store_transaction = AsyncMock(return_value=True)
monkeypatch.setattr(wallet_module, "find_trusted_mint_with_funds", find_funded)
monkeypatch.setattr(wallet_module, "get_wallet", AsyncMock(return_value=wallet))
monkeypatch.setattr(
wallet_module, "get_proofs_per_mint_and_unit", Mock(return_value=proofs)
)
monkeypatch.setattr(admin, "store_cashu_transaction", store_transaction)
result = await admin.withdraw(
Mock(), admin.WithdrawRequest(amount=75, mint_url=requested_mint)
)
assert result == {"token": token, "mint_url": actual_mint}
find_funded.assert_awaited_once_with(
75, "sat", requested_mint, force_reload=True
)
wallet.select_to_send.assert_awaited_once()
store_transaction.assert_awaited_once_with(
token=token,
amount=75,
unit="sat",
mint_url=actual_mint,
typ="out",
collected=False,
source="admin",
)
@pytest.mark.asyncio
async def test_withdraw_maps_true_aggregate_insufficient_funds_to_400(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(
admin,
"send_token",
AsyncMock(
side_effect=ValueError(
"No trusted mint has 75 sat available; balances={'mint': 0}"
)
),
)
with pytest.raises(HTTPException) as exc_info:
await admin.withdraw(Mock(), admin.WithdrawRequest(amount=75))
assert exc_info.value.status_code == 400
assert exc_info.value.detail == "Insufficient wallet balance"
+48
View File
@@ -270,6 +270,54 @@ async def test_internal_error_with_invalid_keyword_does_not_masquerade(
assert await session.get(ApiKey, hashed_key) is None assert await session.get(ApiKey, hashed_key) is None
@pytest.mark.asyncio
async def test_primary_msat_token_sets_provenance_without_cashu_mint_duplicate(
session: AsyncSession,
) -> None:
token = "cashuAprimary_msat_token"
token_obj = SimpleNamespace(mint="http://primary:3338", unit="msat")
credit = AsyncMock(return_value=1_000)
from routstr.core.settings import settings
with (
patch.object(settings, "primary_mint", token_obj.mint),
patch.object(settings, "primary_mint_unit", "msat"),
patch.object(settings, "cashu_mints", []),
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
patch("routstr.auth.credit_balance", new=credit),
):
key = await validate_bearer_key(token, session)
assert key.refund_mint_url == token_obj.mint
assert key.refund_currency == "msat"
credit.assert_awaited_once()
@pytest.mark.asyncio
async def test_primary_token_unit_mismatch_is_rejected_before_redemption(
session: AsyncSession,
) -> None:
token = "cashuAprimary_wrong_unit"
token_obj = SimpleNamespace(mint="http://primary:3338", unit="sat")
credit = AsyncMock(return_value=1_000)
from routstr.core.settings import settings
with (
patch.object(settings, "primary_mint", token_obj.mint),
patch.object(settings, "primary_mint_unit", "msat"),
patch.object(settings, "cashu_mints", []),
patch("routstr.auth.deserialize_token_from_string", return_value=token_obj),
patch("routstr.auth.credit_balance", new=credit),
):
with pytest.raises(HTTPException) as exc_info:
await validate_bearer_key(token, session)
assert exc_info.value.status_code == 400
credit.assert_not_awaited()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_malformed_cashu_token_returns_400_invalid_token( async def test_malformed_cashu_token_returns_400_invalid_token(
session: AsyncSession, session: AsyncSession,
+11 -14
View File
@@ -4,7 +4,7 @@ Tests admin endpoints that are testable without full app setup:
withdraw validation, authentication guards, and slug validation. withdraw validation, authentication guards, and slug validation.
""" """
from unittest.mock import Mock, patch from unittest.mock import AsyncMock, patch
import pytest import pytest
from fastapi import HTTPException, Request from fastapi import HTTPException, Request
@@ -46,22 +46,19 @@ async def test_withdraw_rejects_insufficient_balance() -> None:
request = Request(scope={"type": "http", "method": "POST"}) request = Request(scope={"type": "http", "method": "POST"})
with patch("routstr.core.admin.get_wallet") as mock_wallet, \ with patch(
patch("routstr.core.admin.get_proofs_per_mint_and_unit") as mock_proofs, \ "routstr.core.admin.send_token",
patch("routstr.core.admin.slow_filter_spend_proofs") as mock_filter: new=AsyncMock(
side_effect=ValueError(
mock_w = Mock() "No trusted mint has 1000000 sat available; balances={}"
mock_w.keysets = {} )
mock_w.proofs = [] ),
mock_wallet.return_value = mock_w ):
mock_proofs.return_value = []
mock_filter.return_value = []
with pytest.raises(HTTPException) as exc_info: with pytest.raises(HTTPException) as exc_info:
await withdraw(request, WithdrawRequest(amount=1000000, unit="sat")) await withdraw(request, WithdrawRequest(amount=1000000, unit="sat"))
assert exc_info.value.status_code == 400 assert exc_info.value.status_code == 400
assert "Insufficient" in str(exc_info.value.detail) assert "Insufficient" in str(exc_info.value.detail)
# =========================================================================== # ===========================================================================
+1 -1
View File
@@ -67,7 +67,7 @@ async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> N
payout_wallet = Mock() payout_wallet = Mock()
events: list[str] = [] events: list[str] = []
async def prepare(*_args: object) -> Mock: async def prepare(*_args: object, **_kwargs: object) -> Mock:
events.append("prepare") events.append("prepare")
return payout_wallet return payout_wallet
+158 -4
View File
@@ -6,13 +6,16 @@ from unittest.mock import AsyncMock, Mock, patch
import httpx import httpx
import pytest import pytest
from cashu.core.base import Proof from cashu.core.base import MintQuoteState, Proof
from routstr.lightning import ( from routstr.lightning import (
InvoiceRecoverRequest,
_invoice_settlement_locks, _invoice_settlement_locks,
_is_outputs_already_signed, _is_outputs_already_signed,
_mint_invoice_quote, _mint_invoice_quote,
check_invoice_payment, check_invoice_payment,
get_invoice_status,
recover_invoice,
) )
from routstr.wallet import Wallet from routstr.wallet import Wallet
@@ -30,6 +33,8 @@ def _invoice(**overrides: object) -> SimpleNamespace:
"balance_limit": None, "balance_limit": None,
"balance_limit_reset": None, "balance_limit_reset": None,
"validity_date": None, "validity_date": None,
"created_at": 1,
"expires_at": 2,
} }
values.update(overrides) values.update(overrides)
return SimpleNamespace(**values) return SimpleNamespace(**values)
@@ -159,14 +164,21 @@ async def test_non_pending_invoice_is_not_minted() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ambiguous_invoice_mint_timeout_does_not_expose_paid() -> None: async def test_ambiguous_invoice_mint_timeout_remains_recoverable() -> None:
_invoice_settlement_locks.clear() _invoice_settlement_locks.clear()
invoice = _invoice() invoice = _invoice()
session = AsyncMock() session = AsyncMock()
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True))) wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True)))
state_session = AsyncMock()
state_session.exec.return_value.rowcount = 1
@asynccontextmanager
async def owned_session() -> AsyncIterator[AsyncMock]:
yield state_session
with ( with (
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
patch("routstr.lightning.create_session", owned_session),
patch( patch(
"routstr.lightning._mint_invoice_quote", "routstr.lightning._mint_invoice_quote",
AsyncMock(side_effect=httpx.TimeoutException("response lost")), AsyncMock(side_effect=httpx.TimeoutException("response lost")),
@@ -175,12 +187,152 @@ async def test_ambiguous_invoice_mint_timeout_does_not_expose_paid() -> None:
): ):
await check_invoice_payment(invoice, session) # type: ignore[arg-type] await check_invoice_payment(invoice, session) # type: ignore[arg-type]
assert invoice.status == "pending" assert invoice.status == "settlement_pending"
state_session.commit.assert_awaited_once()
session.rollback.assert_not_awaited() session.rollback.assert_not_awaited()
# One commit closes the initial read transaction before external I/O. # One commit closes the initial read transaction before external I/O.
session.commit.assert_awaited_once() session.commit.assert_awaited_once()
@pytest.mark.asyncio
async def test_quote_lookup_timeout_is_not_definitively_unpaid() -> None:
_invoice_settlement_locks.clear()
invoice = _invoice(expires_at=0)
session = AsyncMock()
wallet = Mock(
get_mint_quote=AsyncMock(side_effect=httpx.TimeoutException("quote timeout"))
)
with (
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
patch("routstr.lightning._reload_invoice_view", AsyncMock()),
):
result = await check_invoice_payment(invoice, session) # type: ignore[arg-type]
assert result is False
@pytest.mark.asyncio
async def test_overdue_invoice_does_not_expire_after_ambiguous_quote_lookup() -> None:
invoice = _invoice(status="pending", expires_at=0)
session = AsyncMock()
session.get.return_value = invoice
check = AsyncMock(return_value=False)
with patch("routstr.lightning.check_invoice_payment", check):
response = await get_invoice_status(invoice.id, session) # type: ignore[arg-type]
assert response.status == "pending"
session.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_overdue_invoice_expires_only_after_definitive_unpaid_quote() -> None:
invoice = _invoice(status="pending", expires_at=0)
session = AsyncMock()
session.get.return_value = invoice
check = AsyncMock(return_value=True)
async def expire(
candidate: SimpleNamespace, _session: AsyncMock, definitive: bool
) -> bool:
assert definitive is True
candidate.status = "expired"
return True
with (
patch("routstr.lightning.check_invoice_payment", check),
patch(
"routstr.lightning._expire_invoice_if_authoritatively_unpaid",
side_effect=expire,
) as expire_invoice,
):
response = await get_invoice_status(invoice.id, session) # type: ignore[arg-type]
assert response.status == "expired"
expire_invoice.assert_awaited_once_with(invoice, session, True)
session.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_recover_applies_authoritative_expiry_helper() -> None:
invoice = _invoice(status="pending", expires_at=0)
session = AsyncMock()
result = Mock()
result.first.return_value = invoice
session.exec.return_value = result
check = AsyncMock(return_value=True)
async def expire(
candidate: SimpleNamespace, _session: AsyncMock, definitive: bool
) -> bool:
assert definitive is True
candidate.status = "expired"
return True
with (
patch("routstr.lightning.check_invoice_payment", check),
patch(
"routstr.lightning._expire_invoice_if_authoritatively_unpaid",
side_effect=expire,
) as expire_invoice,
):
response = await recover_invoice(
InvoiceRecoverRequest(bolt11="lnbc-test"), session # type: ignore[arg-type]
)
assert response.status == "expired"
expire_invoice.assert_awaited_once_with(invoice, session, True)
@pytest.mark.asyncio
async def test_paid_state_write_failure_still_reports_non_expirable_outcome() -> None:
_invoice_settlement_locks.clear()
invoice = _invoice(expires_at=0)
session = AsyncMock()
wallet = Mock(
get_mint_quote=AsyncMock(
return_value=Mock(paid=True, state=MintQuoteState.paid)
)
)
with (
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
patch(
"routstr.lightning._mint_invoice_quote",
AsyncMock(side_effect=httpx.TimeoutException("response lost")),
),
patch(
"routstr.lightning.create_session",
side_effect=RuntimeError("database unavailable"),
),
patch("routstr.lightning._reload_invoice_view", AsyncMock()),
):
definitively_unpaid = await check_invoice_payment(
invoice, session # type: ignore[arg-type]
)
assert definitively_unpaid is False
assert invoice.status == "pending"
@pytest.mark.asyncio
async def test_settlement_pending_invoice_does_not_expire() -> None:
invoice = _invoice(status="settlement_pending", expires_at=0)
session = AsyncMock()
session.get.return_value = invoice
check = AsyncMock()
with patch("routstr.lightning.check_invoice_payment", check):
response = await get_invoice_status(
invoice.id, session # type: ignore[arg-type]
)
check.assert_awaited_once_with(invoice, session)
assert response.status == "settlement_pending"
session.commit.assert_not_awaited()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_concurrent_invoice_checks_finalize_once_in_process() -> None: async def test_concurrent_invoice_checks_finalize_once_in_process() -> None:
_invoice_settlement_locks.clear() _invoice_settlement_locks.clear()
@@ -195,7 +347,9 @@ async def test_concurrent_invoice_checks_finalize_once_in_process() -> None:
@asynccontextmanager @asynccontextmanager
async def owned_session() -> AsyncIterator[AsyncMock]: async def owned_session() -> AsyncIterator[AsyncMock]:
yield AsyncMock() owned = AsyncMock()
owned.exec.return_value.rowcount = 1
yield owned
with ( with (
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
+72
View File
@@ -4,10 +4,12 @@ import asyncio
from typing import Any from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest import pytest
from cashu.core.base import MeltQuoteState from cashu.core.base import MeltQuoteState
from routstr.core.settings import settings from routstr.core.settings import settings
from routstr.mint import MintCooldownError, MintRateGuard
from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl from routstr.payment.lnurl import LNURLError, raw_send_to_lnurl
LNURL_DATA = { LNURL_DATA = {
@@ -111,6 +113,76 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None:
wallet.get_melt_quote.assert_awaited_once_with("q") wallet.get_melt_quote.assert_awaited_once_with("q")
@pytest.mark.asyncio
@pytest.mark.parametrize("rate_error", ["cooldown", "http_429"])
async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs(
rate_error: str,
) -> None:
wallet, proofs = _wallet()
wallet.melt = AsyncMock()
wallet.set_reserved_for_send = AsyncMock()
data_patch, invoice_patch = _lnurl_patches()
async def run_operation(factory: Any, *, op_name: str, **_: object) -> Any:
if op_name == "lnurl_melt":
if rate_error == "cooldown":
raise MintCooldownError(str(wallet.url), 60)
request = httpx.Request("POST", f"{wallet.url}/v1/melt/bolt11")
response = httpx.Response(429, request=request)
raise httpx.HTTPStatusError(
"rate limited", request=request, response=response
)
return await factory()
with (
data_patch,
invoice_patch,
patch(
"routstr.payment.lnurl.run_mint_operation",
side_effect=run_operation,
),
pytest.raises((MintCooldownError, httpx.HTTPStatusError)),
):
await raw_send_to_lnurl(
wallet, proofs, "owner@ln.tld", "sat", amount=1000
)
wallet.melt.assert_not_awaited()
wallet.set_reserved_for_send.assert_awaited_once_with(
proofs, reserved=False
)
@pytest.mark.asyncio
async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None:
wallet, proofs = _wallet()
request = httpx.Request("POST", f"{wallet.url}/v1/melt/bolt11")
response = httpx.Response(429, request=request)
wallet.melt = AsyncMock(
side_effect=httpx.HTTPStatusError(
"rate limited", request=request, response=response
)
)
wallet.set_reserved_for_send = AsyncMock()
data_patch, invoice_patch = _lnurl_patches()
with (
patch.object(settings, "mint_retry_max_attempts", 0),
data_patch,
invoice_patch,
pytest.raises(httpx.HTTPStatusError),
):
await raw_send_to_lnurl(
wallet, proofs, "owner@ln.tld", "sat", amount=1000
)
wallet.melt.assert_awaited_once()
wallet.set_reserved_for_send.assert_awaited_once_with(
proofs, reserved=False
)
MintRateGuard._guards.pop(str(wallet.url), None)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_raw_send_to_lnurl_succeeds_on_explicit_paid_response() -> None: async def test_raw_send_to_lnurl_succeeds_on_explicit_paid_response() -> None:
wallet, proofs = _wallet() wallet, proofs = _wallet()
+56
View File
@@ -1,3 +1,4 @@
import asyncio
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
import httpx import httpx
@@ -31,6 +32,43 @@ async def test_cooldown_fails_fast_while_wallet_mutation_scope_is_held() -> None
sleep.assert_not_awaited() sleep.assert_not_awaited()
@pytest.mark.asyncio
async def test_expired_cooldown_allows_probe_in_wallet_mutation_scope() -> None:
guard = MintRateGuard("http://mint:3338", max_concurrency=1)
guard.apply_cooldown(0, reason="rate_limited")
operation = AsyncMock(return_value="recovered")
async with fail_fast_mint_operations():
result = await guard.run(operation)
assert result == "recovered"
operation.assert_awaited_once()
assert guard._needs_probe is False
@pytest.mark.asyncio
async def test_fail_fast_does_not_wait_behind_existing_probe() -> None:
guard = MintRateGuard("http://mint:3338", max_concurrency=1)
guard.apply_cooldown(0, reason="rate_limited")
probe_started = asyncio.Event()
release_probe = asyncio.Event()
async def probe() -> str:
probe_started.set()
await release_probe.wait()
return "recovered"
first = asyncio.create_task(guard.run(probe))
await probe_started.wait()
try:
async with fail_fast_mint_operations():
with pytest.raises(MintCooldownError):
await asyncio.wait_for(guard.run(AsyncMock()), timeout=0.05)
finally:
release_probe.set()
assert await first == "recovered"
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_cashu_429_dispatches_through_wallet_override() -> None: async def test_cashu_429_dispatches_through_wallet_override() -> None:
async def handler(request: httpx.Request) -> httpx.Response: async def handler(request: httpx.Request) -> httpx.Response:
@@ -63,3 +101,21 @@ async def test_cashu_429_dispatches_through_wallet_override() -> None:
pytest.raises(MintRateLimitedError), pytest.raises(MintRateLimitedError),
): ):
await wallet.mint_quote(1, Unit.sat) await wallet.mint_quote(1, Unit.sat)
async def test_guard_concurrency_change_preserves_cooldown_state() -> None:
from routstr.core.settings import settings
mint_url = "https://mint.test-concurrency-carryover"
with patch.object(settings, "mint_max_concurrency", 2):
guard = MintRateGuard.get(mint_url)
guard.apply_cooldown(120.0, reason="rate_limited")
guard._consecutive_rate_limits = 3
with patch.object(settings, "mint_max_concurrency", 5):
rebuilt = MintRateGuard.get(mint_url)
assert rebuilt is not guard
assert rebuilt.cooldown_remaining() > 0
assert rebuilt._cooldown_reason == "rate_limited"
assert rebuilt._consecutive_rate_limits == 3
+30
View File
@@ -125,3 +125,33 @@ async def test_get_max_cost_for_model_tolerance() -> None:
"gpt-4", session=mock_session, model_obj=mock_model "gpt-4", session=mock_session, model_obj=mock_model
) )
assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000 assert cost == 450000 # 500 sats * 1000 * 0.9 = 450000
async def test_discounted_max_cost_floors_at_min_request_msat() -> None:
from routstr.payment.helpers import calculate_discounted_max_cost
pricing = Mock()
pricing.prompt = 0.001
pricing.completion = 0.001
pricing.max_prompt_cost = 100.0
pricing.max_completion_cost = 100.0
model_obj = Mock()
model_obj.sats_pricing = pricing
model_obj.top_provider = None
model_obj.context_length = None
body = {
"model": "test-model",
"messages": [{"role": "user", "content": "hi"}],
"max_tokens": 1,
}
with (
patch.object(settings, "fixed_pricing", False),
patch.object(settings, "tolerance_percentage", 0),
patch.object(settings, "min_request_msat", 1000),
):
cost = await calculate_discounted_max_cost(150_000, body, model_obj)
assert cost == 1000
+254 -5
View File
@@ -2,7 +2,8 @@ import asyncio
import base64 import base64
import json import json
import socket import socket
from collections.abc import Generator from collections.abc import AsyncIterator, Generator
from contextlib import asynccontextmanager
from unittest.mock import AsyncMock, Mock, patch from unittest.mock import AsyncMock, Mock, patch
import httpx import httpx
@@ -61,6 +62,49 @@ async def test_get_balance() -> None:
assert balance == 50000 assert balance == 50000
@pytest.mark.asyncio
async def test_get_wallet_force_reload_bypasses_reload_interval() -> None:
from routstr.wallet import get_wallet
mock_wallet = Mock(load_mint=AsyncMock(), load_proofs=AsyncMock())
with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)):
await get_wallet("http://mint:3338", "sat")
await get_wallet("http://mint:3338", "sat", force_reload=True)
assert mock_wallet.load_mint.await_count == 2
assert mock_wallet.load_proofs.await_count == 2
@pytest.mark.asyncio
async def test_public_recieve_token_holds_wallet_operation_guard() -> None:
inside_guard = False
@asynccontextmanager
async def operation_guard() -> AsyncIterator[None]:
nonlocal inside_guard
inside_guard = True
try:
yield
finally:
inside_guard = False
async def receive_locked(*_args: object, **_kwargs: object) -> tuple[int, str, str]:
assert inside_guard
return 1, "sat", "https://mint.example"
with (
patch("routstr.wallet.wallet_operation_guard", operation_guard),
patch("routstr.wallet._recieve_token_locked", side_effect=receive_locked),
):
assert await recieve_token("cashuAtoken") == (
1,
"sat",
"https://mint.example",
)
assert inside_guard is False
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_recieve_token_valid() -> None: async def test_recieve_token_valid() -> None:
token_data = { token_data = {
@@ -167,6 +211,78 @@ async def test_recieve_token_trusted_mint_deducts_input_fee() -> None:
) )
@pytest.mark.asyncio
async def test_recieve_token_uses_only_requested_destination_mint() -> None:
from routstr.core.settings import settings
source = "http://foreign:3338"
destination = "http://key-mint:3338"
token = Mock(
mint=source,
unit="sat",
amount=100,
keysets=["keyset1"],
proofs=[Mock(amount=100)],
)
source_wallet = Mock()
swap = AsyncMock(return_value=(99, "sat", destination))
with (
patch.object(settings, "primary_mint", destination),
patch.object(settings, "cashu_mints", [destination]),
patch("routstr.wallet.deserialize_token_from_string", return_value=token),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=source_wallet)),
patch("routstr.wallet.swap_to_trusted_mint", swap),
):
result = await recieve_token(
"cashuAtoken", destination_mint=destination, destination_unit="sat"
)
assert result == (99, "sat", destination)
swap.assert_awaited_once_with(
token, source_wallet, destination_mints=[destination]
)
@pytest.mark.asyncio
async def test_recieve_token_rejects_unit_mismatch_before_wallet_mutation() -> None:
token = Mock(mint="http://key-mint:3338", unit="msat", keysets=["keyset"])
get_wallet = AsyncMock()
with (
patch("routstr.wallet.deserialize_token_from_string", return_value=token),
patch("routstr.wallet.get_wallet", get_wallet),
pytest.raises(ValueError, match="liability unit"),
):
await recieve_token(
"cashuAtoken",
destination_mint="http://key-mint:3338",
destination_unit="sat",
)
get_wallet.assert_not_awaited()
@pytest.mark.asyncio
async def test_recieve_token_cross_mint_output_unit_must_match() -> None:
token = Mock(mint="http://foreign:3338", unit="msat", keysets=["keyset"])
get_wallet = AsyncMock()
with (
patch("routstr.wallet.deserialize_token_from_string", return_value=token),
patch("routstr.wallet.settings.primary_mint_unit", "sat"),
patch("routstr.wallet.get_wallet", get_wallet),
pytest.raises(ValueError, match="liability unit"),
):
await recieve_token(
"cashuAtoken",
destination_mint="http://key-mint:3338",
destination_unit="msat",
)
get_wallet.assert_not_awaited()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_primary_mint_failure_does_not_try_another_mint() -> None: async def test_primary_mint_failure_does_not_try_another_mint() -> None:
from routstr.core.settings import settings from routstr.core.settings import settings
@@ -207,6 +323,54 @@ async def test_primary_mint_failure_does_not_try_another_mint() -> None:
assert failure["action"] == "retry_with_token_from_another_mint" assert failure["action"] == "retry_with_token_from_another_mint"
@pytest.mark.asyncio
async def test_same_mint_split_timeout_is_non_retryable() -> None:
from routstr.wallet import _redeem_same_mint
token = Mock(
keysets=["keyset1"],
mint="http://mint:3338",
unit="sat",
amount=1000,
proofs=[Mock(amount=1000)],
)
wallet = Mock(
load_mint=AsyncMock(),
split=AsyncMock(side_effect=httpx.ReadTimeout("response lost")),
get_fees_for_proofs=Mock(return_value=0),
)
with pytest.raises(TokenConsumedError, match="outcome is ambiguous") as caught:
await _redeem_same_mint(wallet, token)
classified = classify_redemption_error(caught.value)
assert classified is not None
assert classified[0] == "token_consumed"
assert classified[1] == 500
assert classified[3] == "cashu_token_consumed"
@pytest.mark.asyncio
async def test_same_mint_split_connect_error_remains_retryable() -> None:
from routstr.wallet import SourceMintConnectionError, _redeem_same_mint
token = Mock(
keysets=["keyset1"],
mint="http://mint:3338",
unit="sat",
amount=1000,
proofs=[Mock(amount=1000)],
)
wallet = Mock(
load_mint=AsyncMock(),
split=AsyncMock(side_effect=httpx.ConnectError("connect failed")),
get_fees_for_proofs=Mock(return_value=0),
)
with pytest.raises(SourceMintConnectionError):
await _redeem_same_mint(wallet, token)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_send_token() -> None: async def test_send_token() -> None:
mock_wallet = Mock() mock_wallet = Mock()
@@ -224,7 +388,11 @@ async def test_release_token_reservation_unreserves_local_proofs() -> None:
token_proof = Mock(secret="proof-secret", reserved=True) token_proof = Mock(secret="proof-secret", reserved=True)
cached_proof = Mock(secret="proof-secret", reserved=True) cached_proof = Mock(secret="proof-secret", reserved=True)
token = Mock(mint="http://mint:3338", unit="sat", proofs=[token_proof]) token = Mock(mint="http://mint:3338", unit="sat", proofs=[token_proof])
wallet = Mock(proofs=[cached_proof], set_reserved_for_send=AsyncMock()) wallet = Mock(
proofs=[cached_proof],
load_proofs=AsyncMock(),
set_reserved_for_send=AsyncMock(),
)
with ( with (
patch("routstr.wallet.deserialize_token_from_string", return_value=token), patch("routstr.wallet.deserialize_token_from_string", return_value=token),
patch( patch(
@@ -234,6 +402,7 @@ async def test_release_token_reservation_unreserves_local_proofs() -> None:
await release_token_reservation("cashu-token") await release_token_reservation("cashu-token")
get_wallet.assert_awaited_once_with("http://mint:3338", "sat", load=False) get_wallet.assert_awaited_once_with("http://mint:3338", "sat", load=False)
wallet.load_proofs.assert_awaited_once_with(reload=True)
wallet.set_reserved_for_send.assert_awaited_once_with(token.proofs, reserved=False) wallet.set_reserved_for_send.assert_awaited_once_with(token.proofs, reserved=False)
assert token_proof.reserved is False assert token_proof.reserved is False
assert cached_proof.reserved is False assert cached_proof.reserved is False
@@ -270,6 +439,64 @@ async def test_refund_mint_falls_back_to_trusted_mint_with_funds() -> None:
assert mint == secondary assert mint == secondary
@pytest.mark.asyncio
async def test_send_refreshes_reservations_inside_wallet_guard() -> None:
mint = "http://mint:3338"
proof = Mock(amount=1000, reserved=False)
wallet = Mock(
keysets={},
proofs=[proof],
select_to_send=AsyncMock(return_value=([proof], None)),
serialize_proofs=AsyncMock(return_value="token"),
set_reserved_for_send=AsyncMock(),
)
inside_guard = False
@asynccontextmanager
async def operation_guard() -> AsyncIterator[None]:
nonlocal inside_guard
inside_guard = True
try:
yield
finally:
inside_guard = False
async def find_mint(
amount: int,
unit: str,
preferred_mint: str | None,
*,
force_reload: bool,
) -> str:
assert inside_guard
assert (amount, unit, preferred_mint, force_reload) == (
1000,
"sat",
mint,
True,
)
return mint
async def get_loaded_wallet(*_: object, **__: object) -> Mock:
assert inside_guard
return wallet
with (
patch("routstr.wallet.wallet_operation_guard", operation_guard),
patch("routstr.wallet.find_trusted_mint_with_funds", side_effect=find_mint),
patch("routstr.wallet.get_wallet", side_effect=get_loaded_wallet),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
return_value=[proof],
),
):
assert await send(1000, "sat", mint) == (1000, "token")
wallet.set_reserved_for_send.assert_awaited_once_with(
[proof], reserved=True
)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_send_falls_back_when_preferred_mint_has_only_reserved_balance() -> None: async def test_send_falls_back_when_preferred_mint_has_only_reserved_balance() -> None:
from routstr.core.settings import settings from routstr.core.settings import settings
@@ -395,6 +622,27 @@ async def test_credit_balance() -> None:
assert mock_session.refresh.called assert mock_session.refresh.called
@pytest.mark.asyncio
async def test_credit_balance_constrains_redemption_to_key_mint() -> None:
key_mint = "http://key-mint:3338"
mock_key = Mock(
balance=1_000_000,
hashed_key="test_hash",
refund_mint_url=key_mint,
refund_currency="sat",
)
mock_session = AsyncMock()
mock_session.exec.return_value.rowcount = 1
receive = AsyncMock(return_value=(1000, "sat", key_mint))
with patch("routstr.wallet.recieve_token", receive):
await credit_balance("cashuAtoken", mock_key, mock_session)
receive.assert_awaited_once_with(
"cashuAtoken", destination_mint=key_mint, destination_unit="sat"
)
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_credit_balance_rejects_zero_amount() -> None: async def test_credit_balance_rejects_zero_amount() -> None:
"""A zero/dust redemption must raise BEFORE any commit, so no orphan """A zero/dust redemption must raise BEFORE any commit, so no orphan
@@ -2386,12 +2634,12 @@ def test_classify_500_with_rate_limit_text_is_not_mint_rate_limited() -> None:
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# _MintRateGuard — probe does NOT escalate cooldown counter # _MintRateGuard — probe backoff escalation and recovery
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_probe_does_not_escalate_consecutive_rate_limits() -> None: async def test_probe_escalates_consecutive_rate_limits() -> None:
from routstr.mint import MintRateGuard from routstr.mint import MintRateGuard
guard = MintRateGuard("http://mint", max_concurrency=0) guard = MintRateGuard("http://mint", max_concurrency=0)
@@ -2401,8 +2649,9 @@ async def test_probe_does_not_escalate_consecutive_rate_limits() -> None:
with pytest.raises(httpx.HTTPStatusError): with pytest.raises(httpx.HTTPStatusError):
await guard.run(AsyncMock(side_effect=_http_429_error())) await guard.run(AsyncMock(side_effect=_http_429_error()))
assert guard._consecutive_rate_limits == 1 assert guard._consecutive_rate_limits == 2
assert guard._needs_probe is True assert guard._needs_probe is True
assert guard.cooldown_remaining() > 60
@pytest.mark.asyncio @pytest.mark.asyncio
+1
View File
@@ -42,6 +42,7 @@ export interface BalanceDetail {
export interface WithdrawResponse { export interface WithdrawResponse {
token: string; token: string;
mint_url: string;
} }
export interface CreateChildKeyResponse { export interface CreateChildKeyResponse {