diff --git a/.env.example b/.env.example index 35a171ba..265ca2c4 100644 --- a/.env.example +++ b/.env.example @@ -33,6 +33,8 @@ ROUTSTR_SECRET_KEY= # DATABASE_POOL_PRE_PING=false # Warn when a checkout is held this many seconds. # DATABASE_POOL_HOLD_WARN_SECONDS=10 +# SQLite write-lock timeout, in seconds. +# DATABASE_BUSY_TIMEOUT=30 # SQLite serialises writes; increasing its pool can trade pool timeouts for # "database is locked" errors rather than increasing write throughput. diff --git a/routstr/auth.py b/routstr/auth.py index 0dfa38e4..7b808007 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -510,11 +510,25 @@ async def _validate_bearer_key_locked( "AUTH: credit_balance returned successfully", extra={"msats": msats} ) except Exception as credit_error: - logger.error( + classification = classify_redemption_error(credit_error) + expected_codes = { + "cashu_token_already_spent", + "cashu_source_mint_unreachable", + "cashu_mint_unreachable", + "cashu_mint_rate_limited", + } + log = ( + logger.info + if classification is not None + and classification[3] in expected_codes + else logger.error + ) + log( "AUTH: credit_balance failed", extra={ "error": str(credit_error), "error_type": type(credit_error).__name__, + "error_code": classification[3] if classification else None, }, ) await session.rollback() @@ -756,13 +770,19 @@ async def pay_for_request( result = await session.exec(stmt) # type: ignore[call-overload] if result.rowcount == 0: - logger.error( - "Concurrent request depleted balance", + await session.refresh(billing_key) + total_balance = billing_key.balance + reserved_balance = billing_key.reserved_balance + available_balance = max(0, total_balance - reserved_balance) + logger.warning( + "Concurrent request depleted available balance", extra={ "key_hash": key.hashed_key[:8] + "...", "billing_key_hash": billing_key.hashed_key[:8] + "...", "required_cost": cost_per_request, - "current_balance": billing_key.balance, + "total_balance": total_balance, + "reserved_balance": reserved_balance, + "available_balance": available_balance, }, ) @@ -770,9 +790,10 @@ async def pay_for_request( status_code=402, detail={ "error": { - "message": f"Insufficient balance: {cost_per_request} mSats required. {billing_key.balance} available.", + "message": f"Insufficient balance: {cost_per_request} mSats required. {available_balance} available.", "type": "insufficient_quota", "code": "insufficient_balance", + "available_balance": available_balance, } }, ) diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 63fb44f6..66e3a59a 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,4 +1,3 @@ -import asyncio import json import re import secrets @@ -1415,50 +1414,37 @@ async def initiate_provider_topup( else {} ) - last_status_code = 500 - last_error_detail: object = "Failed to create top-up invoice" + # Quote creation is unsafe to retry without idempotency. + resp = await client.post( + f"{clean_url}/v1/balance/lightning/invoice", + json=request_json, + headers=headers, + ) - # Some upstream Routstr nodes fail the first invoice request after warm-up - # and succeed immediately on retry. Retry once here so the UI stays single-click. - for attempt in range(2): - resp = await client.post( - f"{clean_url}/v1/balance/lightning/invoice", - json=request_json, - headers=headers, - ) - - if resp.status_code == 200: - data = resp.json() - return { - "ok": True, - "topup_data": { - "payment_request": data.get("bolt11"), - "invoice_id": data.get("invoice_id"), - "status": "pending", - }, - } - - logger.error( - f"Upstream topup request failed: {resp.text}", - extra={ - "provider_id": provider_id, - "attempt": attempt + 1, - "status_code": resp.status_code, + if resp.status_code == 200: + data = resp.json() + return { + "ok": True, + "topup_data": { + "payment_request": data.get("bolt11"), + "invoice_id": data.get("invoice_id"), + "status": "pending", }, - ) - try: - last_error_detail = resp.json() - except Exception: - last_error_detail = resp.text - last_status_code = resp.status_code - - if resp.status_code < 500 or attempt == 1: - break - - await asyncio.sleep(0.2) + } + logger.error( + f"Upstream topup request failed: {resp.text}", + extra={ + "provider_id": provider_id, + "status_code": resp.status_code, + }, + ) + try: + error_detail: object = resp.json() + except Exception: + error_detail = resp.text raise HTTPException( - status_code=last_status_code, detail=last_error_detail + status_code=resp.status_code, detail=error_detail ) upstream_instance = _instantiate_provider(provider) diff --git a/routstr/core/db.py b/routstr/core/db.py index 2ab5ef25..14cc669e 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -37,6 +37,9 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"} pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping} + connect_args: dict[str, object] = {} + if is_sqlite and not is_memory_sqlite: + connect_args["timeout"] = settings.database_busy_timeout if not is_memory_sqlite: options.update( pool_size=settings.database_pool_size, @@ -51,9 +54,12 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine: "database_url_backend": backend, "in_memory_sqlite": is_memory_sqlite, **options, + "connect_args": connect_args, }, ) - created_engine = create_async_engine(database_url, echo=False, **options) + created_engine = create_async_engine( + database_url, echo=False, connect_args=connect_args, **options + ) hold_warn_seconds = settings.database_pool_hold_warn_seconds def record_pool_checkout( diff --git a/routstr/core/exceptions.py b/routstr/core/exceptions.py index fd3cfefa..360b810d 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -44,15 +44,25 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon path = request.url.path # 4xx is client behaviour; the uvicorn access log already records it. - # Only 5xx warrants a server-side warning/error log here. if status_code >= 500: - logger.error( + error_type = None + if isinstance(detail, dict): + error = detail.get("error") + if isinstance(error, dict): + error_type = error.get("type") + log = ( + logger.warning + if error_type in {"mint_unreachable", "mint_rate_limited"} + else logger.error + ) + log( f"HTTP {status_code} on {path}: {detail}", extra={ "request_id": request_id, "status_code": status_code, "detail": detail, "path": path, + "error_type": error_type, }, ) diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 53e71a74..7190f058 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -148,6 +148,9 @@ class Settings(BaseSettings): database_pool_hold_warn_seconds: float = Field( default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS" ) + database_busy_timeout: float = Field( + default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT" + ) # Logging log_level: str = Field(default="INFO", env="LOG_LEVEL") @@ -209,6 +212,7 @@ ENV_ONLY_FIELDS = frozenset( "database_pool_recycle", "database_pool_pre_ping", "database_pool_hold_warn_seconds", + "database_busy_timeout", } ) diff --git a/routstr/lightning.py b/routstr/lightning.py index df48756c..13c6b321 100644 --- a/routstr/lightning.py +++ b/routstr/lightning.py @@ -215,11 +215,18 @@ async def _request_mint_with_fallback( ) continue try: - wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False) + wallet = await get_wallet( + mint_url, + "sat", + retry_on_rate_limit=False, + load_proofs=False, + ) quote = await run_mint_operation( lambda: wallet.request_mint(amount_sats), op_name="request_mint_invoice", mint_url=mint_url, + # Quote creation is unsafe to retry without idempotency. + retry_timeouts=False, retry_on_rate_limit=False, ) return quote.request, quote.quote, mint_url @@ -471,7 +478,7 @@ async def check_invoice_payment( await session.commit() mint_url = settlement.mint_url or settings.primary_mint - wallet = await get_wallet(mint_url, "sat") + wallet = await get_wallet(mint_url, "sat", load_proofs=False) try: mint_status = await run_mint_operation( lambda: wallet.get_mint_quote(settlement.payment_hash), diff --git a/routstr/mint.py b/routstr/mint.py index 64632b9e..e1b60ca2 100644 --- a/routstr/mint.py +++ b/routstr/mint.py @@ -177,6 +177,8 @@ class MintRateGuard: if isinstance(error, httpx.HTTPStatusError): retry_after = parse_retry_after(error.response.headers) self.apply_rate_limit_cooldown(retry_after) + elif is_mint_transport_error(error): + self.apply_cooldown(MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport") else: self.apply_cooldown(1.0) logger.warning( @@ -236,6 +238,17 @@ def mint_cooldown_reason(mint_url: str) -> str | None: return MintRateGuard.get(mint_url).cooldown_reason() +def is_mint_transport_error(error: BaseException) -> bool: + current: BaseException | None = error + seen: set[int] = set() + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, MINT_TRANSPORT_EXCEPTIONS): + return True + current = current.__cause__ or current.__context__ + return False + + def is_mint_rate_limited(error: BaseException) -> bool: """Return whether an exception chain represents HTTP 429/cooldown.""" @@ -269,6 +282,7 @@ async def run_mint_operation( mint_url: str = "", retry_timeouts: bool = True, retry_on_rate_limit: bool = True, + allow_during_cooldown: bool = False, ) -> Any: """Run one mint operation with bounded concurrency and adaptive cooldown.""" @@ -282,7 +296,7 @@ async def run_mint_operation( return await factory() async def invoke() -> Any: - if guard is not None: + if guard is not None and not allow_during_cooldown: return await guard.run(timed_factory) return await timed_factory() @@ -293,6 +307,7 @@ async def run_mint_operation( raise except (asyncio.TimeoutError, httpx.TimeoutException) as exc: if retry_timeouts and attempt < max_attempts - 1: + # Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait. backoff = (2**attempt) + (time.monotonic() % 1.0) logger.warning( "Mint operation timed out, retrying", @@ -305,11 +320,19 @@ async def run_mint_operation( ) await asyncio.sleep(backoff) continue + if guard is not None: + guard.apply_cooldown( + MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport" + ) raise httpx.TimeoutException( f"{op_name} timed out (attempts: {attempt + 1})" ) from exc except Exception as exc: if not is_mint_rate_limited(exc): + if guard is not None and is_mint_transport_error(exc): + guard.apply_cooldown( + MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport" + ) raise backoff = (2**attempt) + (time.monotonic() % 1.0) diff --git a/routstr/payment/lnurl.py b/routstr/payment/lnurl.py index c9f48253..def5bca7 100644 --- a/routstr/payment/lnurl.py +++ b/routstr/payment/lnurl.py @@ -1,7 +1,6 @@ from __future__ import annotations import ipaddress -import math from collections.abc import Awaitable, Callable from typing import Any, TypedDict @@ -10,8 +9,8 @@ from cashu.core.base import MeltQuoteState from cashu.wallet.wallet import Proof, Wallet from ..mint import ( - MINT_TRANSPORT_EXCEPTIONS, is_mint_rate_limited, + is_mint_transport_error, run_mint_operation, ) @@ -226,6 +225,38 @@ async def get_lnurl_invoice( return invoice_data["pr"], invoice_data +def _select_melt_proofs( + wallet: Wallet, + proofs: list[Proof], + *, + quote_amount: int, + fee_reserve: int, + gross_budget: int, +) -> tuple[list[Proof] | None, int]: + """Select proofs that cover the quote and exact NUT-02 input fees. + + Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a + wallet's full balance. Melts accept overpayment and return change, so a + bounded, largest-first selection is both safer and minimizes input fees. + """ + selected: list[Proof] = [] + selected_amount = 0 + required = quote_amount + fee_reserve + for proof in sorted(proofs, key=lambda item: item.amount, reverse=True): + if getattr(proof, "reserved", False) is True: + continue + selected.append(proof) + selected_amount += proof.amount + input_fees = int(wallet.get_fees_for_proofs(selected)) + required = quote_amount + fee_reserve + input_fees + if selected_amount >= required: + if required <= gross_budget: + return selected, 0 + # Covered but over budget; more proofs only raise input fees. + break + return None, max(1, required - min(selected_amount, gross_budget)) + + async def raw_send_to_lnurl( wallet: Wallet, proofs: list[Proof], @@ -281,36 +312,51 @@ async def raw_send_to_lnurl( f"({min_sendable_sat} - {max_sendable_sat} {unit})" ) - estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1 - estimated_fees_msat = estimated_fees_sat * 1000 - final_amount = amount_msat - estimated_fees_msat + final_amount = amount_msat - bolt11_invoice, _ = await get_lnurl_invoice( - lnurl_data["callback_url"], final_amount - ) - - melt_quote_resp = await run_mint_operation( - lambda: wallet.melt_quote(invoice=bolt11_invoice), - op_name="lnurl_melt_quote", - mint_url=str(wallet.url), - ) - - # The invoice comes from the LNURL service, so its amount is untrusted. The - # melt quote is the mint's own reading of it, and it must match what we - # asked to send. Checked before the checkpoint and before reserving, so a - # mismatch leaves no durable state and no locked proofs behind. - quoted_amount = int(melt_quote_resp.amount) - expected_amount = final_amount // 1000 if unit == "sat" else final_amount - if quoted_amount != expected_amount: - raise LNURLError( - f"LNURL invoice amount does not match the requested amount " - f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})" + selected_proofs: list[Proof] | None = None + # Find the largest amount covered by the budget after reserve and input fees. + for _ in range(8): + if final_amount < lnurl_data["min_sendable"]: + raise LNURLError("Cashu melt fees leave no payable LNURL amount") + bolt11_invoice, _ = await get_lnurl_invoice( + lnurl_data["callback_url"], final_amount ) + melt_quote_resp = await run_mint_operation( + lambda: wallet.melt_quote(invoice=bolt11_invoice), + op_name="lnurl_melt_quote", + mint_url=str(wallet.url), + # Quote creation is unsafe to retry without idempotency. + retry_timeouts=False, + ) + + quoted_amount = int(melt_quote_resp.amount) + expected_amount = final_amount // 1000 if unit == "sat" else final_amount + if quoted_amount != expected_amount: + raise LNURLError( + f"LNURL invoice amount does not match the requested amount " + f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})" + ) + + selected_proofs, shortfall = _select_melt_proofs( + wallet, + proofs, + quote_amount=quoted_amount, + fee_reserve=int(melt_quote_resp.fee_reserve), + gross_budget=amount, + ) + if selected_proofs is not None: + break + final_amount -= shortfall * (1000 if unit == "sat" else 1) + else: + raise LNURLError("Cashu melt fees exceed the requested gross amount") if on_melt_quote is not None: await on_melt_quote(melt_quote_resp.quote) - proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True) + assert selected_proofs is not None + proofs = selected_proofs + await wallet.set_reserved_for_send(proofs, reserved=True) try: melt_response = await run_mint_operation( @@ -331,15 +377,29 @@ async def raw_send_to_lnurl( # 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): + if not is_mint_transport_error(error): raise + # Cashu clears reservations on transport errors despite an unknown outcome. + try: + await wallet.set_reserved_for_melt( + proofs, reserved=True, quote_id=melt_quote_resp.quote + ) + except Exception as reservation_error: + raise MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous and its proof reservation could not " + "be restored; proofs must not be retried" + ) from reservation_error melt_response = None melt_error: BaseException | None = error else: melt_error = None - if getattr(melt_response, "state", None) == MeltQuoteState.paid: + melt_state = getattr(melt_response, "state", None) + if melt_state == MeltQuoteState.paid: return final_amount + if melt_state == MeltQuoteState.unpaid: + await wallet.set_reserved_for_send(proofs, reserved=False) + raise LNURLError("Cashu mint confirmed that the melt was unpaid") try: quote = await run_mint_operation( @@ -347,6 +407,8 @@ async def raw_send_to_lnurl( op_name="reconcile_lnurl_melt_quote", mint_url=str(wallet.url), retry_timeouts=False, + # Reconciliation must bypass the cooldown opened by this failure. + allow_during_cooldown=True, ) except Exception as reconciliation_error: raise MeltOutcomeAmbiguousError( @@ -356,6 +418,20 @@ async def raw_send_to_lnurl( if quote is not None and quote.state == MeltQuoteState.paid: return final_amount + if quote is not None and quote.state == MeltQuoteState.unpaid: + # A just-dispatched quote can briefly report unpaid before transitioning. + try: + await wallet.set_reserved_for_melt( + proofs, reserved=True, quote_id=melt_quote_resp.quote + ) + except Exception as reservation_error: + raise MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous and its proof reservation could not " + "be restored; proofs must not be retried" + ) from reservation_error + raise MeltOutcomeAmbiguousError( + "Melt outcome is ambiguous; an immediate unpaid state is not final" + ) from melt_error state = getattr(getattr(quote, "state", None), "value", "unknown") raise MeltOutcomeAmbiguousError( diff --git a/routstr/upstream/auto_topup.py b/routstr/upstream/auto_topup.py index b4325e24..a73bdecf 100644 --- a/routstr/upstream/auto_topup.py +++ b/routstr/upstream/auto_topup.py @@ -15,9 +15,6 @@ from ..core.db import ( UpstreamProviderRow, create_session, ) -from ..core.db import ( - store_cashu_transaction_with_retry as store_cashu_transaction, -) from ..payment.price import sats_usd_price from ..wallet import ( Bolt11PaymentAmbiguous, @@ -27,7 +24,7 @@ from ..wallet import ( maximum_owner_cashu_balance_sats, prepare_bolt11_payment, release_token_reservation, - send_token, + send_token_from_owner_locked, token_mint_url, wallet_operation_guard, ) @@ -49,6 +46,7 @@ PPQ_PHASES = frozenset({PPQ_PHASE_CLAIMED, PPQ_PHASE_IN_FLIGHT, PPQ_PHASE_RECONC PPQ_SETTLEMENT_ATTEMPTS = 5 PPQ_SETTLEMENT_POLL_SECONDS = 2 PPQ_PENDING_TTL_SECONDS = 15 * 60 +PPQ_SETTLED_COOLDOWN_SECONDS = 5 * 60 PPQ_MAX_INVOICE_PREMIUM = 1.10 PPQ_MIN_TOPUP_USD = 1 PPQ_MAX_TOPUP_USD = 500 @@ -99,7 +97,7 @@ async def periodic_auto_topup() -> None: except Exception as e: logger.error( "Auto top-up cycle failed", - extra={"error": str(e), "error_type": type(e).__name__}, + extra={"error": repr(e), "error_type": type(e).__name__}, ) await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS) @@ -130,7 +128,7 @@ async def _run_auto_topup_cycle() -> None: extra={ "provider_id": row.id, "base_url": row.base_url, - "error": str(e), + "error": repr(e), "error_type": type(e).__name__, }, ) @@ -161,7 +159,11 @@ async def _reconcile_all_ppq_claims() -> set[int]: except Exception as e: logger.error( "PPQ claim reconciliation failed", - extra={"provider_id": row.id, "error": str(e)}, + extra={ + "provider_id": row.id, + "error": repr(e), + "error_type": type(e).__name__, + }, ) return active_provider_ids @@ -363,67 +365,59 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) try: - token = await send_token(amount, "sat", mint_url) + async with wallet_operation_guard(): + # Keep the spend cap and audit mutation in one wallet lock. + spent_24h_sats = await _routstr_spent_last_24h_sats() + if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS: + raise ValueError("Routstr auto top-up daily spend cap reached") + token = await send_token_from_owner_locked(amount, "sat", mint_url) + actual_mint_url = token_mint_url(token, mint_url) + try: + await _persist_routstr_token_and_mark_sent( + row, + operation_id, + expected_sats=expected_sats, + token=token, + amount=amount, + mint_url=actual_mint_url, + ) + except Exception: + logger.critical( + "Aborting auto top-up because its token and sent claim " + "could not be persisted atomically", + extra={"provider_id": row.id, "mint_url": actual_mint_url}, + ) + try: + await release_token_reservation(token) + except Exception as error: + logger.critical( + "Failed to release untracked auto-topup token", + extra={ + "provider_id": row.id, + "mint_url": actual_mint_url, + "error": repr(error), + }, + ) + else: + logger.warning( + "Auto-topup token was released after persistence failed", + extra={"provider_id": row.id, "mint_url": actual_mint_url}, + ) + raise except Exception as e: - logger.error( - "Failed to create cashu token for auto top-up", + logger.warning( + "Failed to create or persist cashu token for auto top-up", extra={ "provider_id": row.id, "amount": amount, "mint_url": mint_url, - "error": str(e), + "error": repr(e), + "error_type": type(e).__name__, }, ) await _release_routstr_claim(row, operation_id) return - actual_mint_url = token_mint_url(token, mint_url) - try: - await store_cashu_transaction( - token=token, - amount=amount, - unit="sat", - mint_url=actual_mint_url, - typ="out", - collected=False, - source="auto_topup", - ) - except Exception: - logger.critical( - "Aborting auto top-up because its cashu token could not be persisted", - extra={"provider_id": row.id, "mint_url": actual_mint_url}, - ) - try: - await release_token_reservation(token) - except Exception as error: - logger.critical( - "Failed to release untracked auto-topup token", - extra={ - "provider_id": row.id, - "mint_url": actual_mint_url, - "error": str(error), - }, - ) - else: - logger.warning( - "Auto-topup token was released after persistence failed", - extra={"provider_id": row.id, "mint_url": actual_mint_url}, - ) - await _release_routstr_claim(row, operation_id) - return - - # Move the claim before the network call, not after: a worker that dies - # mid-request must leave behind a claim that says a token may already be - # with the peer. - await _mark_routstr_sent( - row, - operation_id, - expected_sats=expected_sats, - token=token, - amount=amount, - mint_url=actual_mint_url, - ) - result = await provider.topup(token) if "error" in result: @@ -705,7 +699,7 @@ async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) -> ) -async def _mark_routstr_sent( +async def _persist_routstr_token_and_mark_sent( row: UpstreamProviderRow, operation_id: str, *, @@ -714,20 +708,59 @@ async def _mark_routstr_sent( amount: int, mint_url: str, ) -> None: - claim = await _current_routstr_claim(row) - failures = claim.failures if claim else 0 - if not await _advance_routstr_claim( - row, - operation_id, - deadline=int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS, - phase=ROUTSTR_PHASE_SENT, - expected_sats=expected_sats, - failures=failures, - token=token, - amount=amount, - mint_url=mint_url, - ): - raise RuntimeError("Routstr auto top-up claim ownership was lost") + state_id = _routstr_state_id(row) + async with create_session() as session: + state = await session.get(CashuTransaction, state_id) + claim = _parse_routstr_request_id(state.request_id if state else None) + if ( + state is None + or state.collected + or state.swept + or claim is None + or claim.operation_id != operation_id + or claim.phase != ROUTSTR_PHASE_CLAIMED + ): + raise RuntimeError("Routstr auto top-up claim ownership was lost") + + result = await session.exec( # type: ignore[call-overload] + update(CashuTransaction) + .where( + col(CashuTransaction.id) == state_id, + col(CashuTransaction.request_id) == state.request_id, + col(CashuTransaction.collected) == False, # noqa: E712 + col(CashuTransaction.swept) == False, # noqa: E712 + ) + .values( + request_id=_routstr_request_id( + operation_id, + int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS, + ROUTSTR_PHASE_SENT, + expected_sats, + claim.failures, + ), + token=token, + amount=amount, + unit="sat", + mint_url=mint_url, + ) + ) + if (getattr(result, "rowcount", 0) or 0) != 1: + await session.rollback() + raise RuntimeError("Routstr auto top-up claim ownership was lost") + + session.add( + CashuTransaction( + id=uuid.uuid4().hex, + token=token, + amount=amount, + unit="sat", + mint_url=mint_url, + type="out", + collected=False, + source="auto_topup", + ) + ) + await session.commit() async def _current_routstr_claim(row: UpstreamProviderRow) -> RoutstrClaim | None: @@ -1081,7 +1114,11 @@ async def _set_ppq_state_terminal( col(CashuTransaction.collected) == False, # noqa: E712 col(CashuTransaction.swept) == False, # noqa: E712 ) - .values(collected=collected, swept=swept) + .values( + collected=collected, + swept=swept, + created_at=int(time.time()) if collected else CashuTransaction.created_at, + ) ) updated = (getattr(result, "rowcount", 0) or 0) == 1 if updated: @@ -1106,8 +1143,10 @@ async def _reconcile_ppq_state( """ async with create_session() as session: transaction = await session.get(CashuTransaction, _ppq_state_id(row)) - if transaction is None or transaction.collected or transaction.swept: + if transaction is None or transaction.swept: return False + if transaction.collected: + return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS claim = _parse_ppq_request_id(transaction.request_id) if claim is None: @@ -1173,7 +1212,7 @@ async def _reconcile_ppq_state( async def _ppq_provider_is_claimable( - session: AsyncSession, provider_id: int | None + session: AsyncSession, row: UpstreamProviderRow ) -> bool: """Re-read the provider inside the claim transaction. @@ -1184,10 +1223,16 @@ async def _ppq_provider_is_claimable( this the worker could create a claim for a provider that no longer exists, orphaning it forever. """ - if provider_id is None: + if row.id is None: return False - current = await session.get(UpstreamProviderRow, provider_id) - return current is not None and current.provider_type == "ppqai" + current = await session.get(UpstreamProviderRow, row.id) + return bool( + current is not None + and current.enabled + and current.provider_type == "ppqai" + and current.api_key == row.api_key + and current.provider_settings == row.provider_settings + ) async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None: @@ -1198,10 +1243,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None: request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending") async with create_session() as session: - if not await _ppq_provider_is_claimable(session, row.id): + if not await _ppq_provider_is_claimable(session, row): return None existing = await session.get(CashuTransaction, state_id) if existing is not None: + if ( + existing.collected + and not existing.swept + and int(time.time()) - existing.created_at + < PPQ_SETTLED_COOLDOWN_SECONDS + ): + return None result = await session.exec( # type: ignore[call-overload] update(CashuTransaction) .where( @@ -1232,7 +1284,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None: async with create_session() as session: # Same fencing as the update path: the provider must still exist # inside the transaction that creates the claim. - if not await _ppq_provider_is_claimable(session, row.id): + if not await _ppq_provider_is_claimable(session, row): return None session.add( CashuTransaction( @@ -1453,6 +1505,26 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None if balance >= threshold_usd: return + # Require two low-balance reads before creating an invoice. + confirmed_balance = await provider.get_balance() + if ( + confirmed_balance is None + or not math.isfinite(confirmed_balance) + or confirmed_balance < 0 + or confirmed_balance >= threshold_usd + ): + logger.info( + "PPQ auto top-up aborted by balance confirmation", + extra={ + "provider_id": row.id, + "first_balance_usd": balance, + "confirmed_balance_usd": confirmed_balance, + "threshold_usd": threshold_usd, + }, + ) + return + balance = confirmed_balance + # Perform local pricing and owner-funds checks before asking PPQ to create # an invoice. The exact mint quote still has to be checked afterward, but # predictable local failures should not leave abandoned PPQ invoices. diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 4080be0f..8a8b06ea 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1,12 +1,13 @@ from __future__ import annotations import asyncio +import inspect import json import math import traceback import typing import uuid -from collections.abc import AsyncGenerator, AsyncIterator, Iterator +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterator from typing import Any, Mapping, Self, cast import httpx @@ -70,6 +71,32 @@ if typing.TYPE_CHECKING: logger = get_logger(__name__) +async def _aclose_if_needed(resource: object | None) -> None: + if resource is None: + return + close = getattr(resource, "aclose", None) + if close is None: + return + result = close() + if inspect.isawaitable(result): + await result + + +async def _finalize_and_close_stream( + finalize: Callable[[], Awaitable[None]] | None, + response: object | None, + client: httpx.AsyncClient | None, +) -> None: + try: + if finalize is not None: + await finalize() + finally: + try: + await _aclose_if_needed(response) + finally: + await _aclose_if_needed(client) + + CostMetadata = CostData | MaxCostData | dict[str, Any] @@ -993,6 +1020,7 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + client: httpx.AsyncClient | None = None, ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -1034,23 +1062,37 @@ class BaseUpstreamProvider: nonlocal usage_finalized if usage_finalized: return - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return - try: - await adjust_payment_for_tokens( - fresh_key, - {"model": last_model_seen or "unknown", "usage": None}, - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - pass + try: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return + try: + await adjust_payment_for_tokens( + fresh_key, + {"model": last_model_seen or "unknown", "usage": None}, + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + except Exception: + logger.exception( + "Fallback stream billing finalization failed; releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, new_session, reservation_snapshot + ) + ) + except Exception: + logger.exception( + "Fallback stream billing recovery could not access the database", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) def _process_event( raw_event: bytes, final: bool = False @@ -1280,18 +1322,24 @@ class BaseUpstreamProvider: except Exception as stream_error: logger.warning( - "Streaming interrupted; finalizing in background", + "Streaming interrupted; finalizing before closing upstream", extra={ "error": str(stream_error), + "error_type": type(stream_error).__name__, "key_hash": key.hashed_key[:8] + "...", }, ) raise finally: - if not usage_finalized: - # Create a background task to ensure finalization happens - # even if the generator is closed early - background_tasks.add_task(finalize_db_only) + # Shielded so a client disconnect cannot cancel billing + # finalization or leak the upstream connection. + await asyncio.shield( + _finalize_and_close_stream( + None if usage_finalized else finalize_db_only, + response, + client, + ) + ) # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -1451,6 +1499,7 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, + client: httpx.AsyncClient | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -1484,23 +1533,37 @@ class BaseUpstreamProvider: nonlocal usage_finalized if usage_finalized: return - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return - try: - await adjust_payment_for_tokens( - fresh_key, - {"model": last_model_seen or "unknown", "usage": None}, - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - pass + try: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return + try: + await adjust_payment_for_tokens( + fresh_key, + {"model": last_model_seen or "unknown", "usage": None}, + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + except Exception: + logger.exception( + "Fallback Responses billing finalization failed; releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, new_session, reservation_snapshot + ) + ) + except Exception: + logger.exception( + "Fallback Responses billing recovery could not access the database", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) def _process_event( raw_event: bytes, final: bool = False @@ -1690,16 +1753,24 @@ class BaseUpstreamProvider: except Exception as stream_error: logger.warning( - "Responses API streaming interrupted; finalizing in background", + "Responses API streaming interrupted; finalizing before closing upstream", extra={ "error": str(stream_error), + "error_type": type(stream_error).__name__, "key_hash": key.hashed_key[:8] + "...", }, ) raise finally: - if not usage_finalized: - await finalize_db_only() + # Shielded so a client disconnect cannot cancel billing + # finalization or leak the upstream connection. + await asyncio.shield( + _finalize_and_close_stream( + None if usage_finalized else finalize_db_only, + response, + client, + ) + ) # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -3051,9 +3122,7 @@ class BaseUpstreamProvider: if is_streaming and response.status_code == 200: background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result = await self.handle_streaming_chat_completion( + return await self.handle_streaming_chat_completion( response, key, max_cost_for_model, @@ -3061,9 +3130,8 @@ class BaseUpstreamProvider: requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + client=client, ) - result.background = background_tasks - return result # Handle both non-streaming chat completions and embeddings if response.status_code == 200: @@ -3332,19 +3400,15 @@ class BaseUpstreamProvider: ) if is_streaming and response.status_code == 200: - result = await self.handle_streaming_responses_completion( + return await self.handle_streaming_responses_completion( response, key, max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, + client=client, ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result.background = background_tasks - return result if response.status_code == 200: try: @@ -3628,54 +3692,30 @@ class BaseUpstreamProvider: extra={"amount": amount, "unit": unit, "mint": mint}, ) - max_retries = 3 - last_exception = None - refund_token = None - - for attempt in range(max_retries): - try: - refund_token = await send_token(amount, unit=unit, mint_url=mint) - break - except Exception as e: - last_exception = e - if attempt < max_retries - 1: - logger.warning( - "Refund token creation failed, retrying", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "attempt": attempt + 1, - "max_retries": max_retries, - "amount": amount, - "unit": unit, - "mint": mint, - }, - ) - else: - logger.error( - "Failed to create refund token after all retries", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "attempt": attempt + 1, - "max_retries": max_retries, - "amount": amount, - "unit": unit, - "mint": mint, - }, - ) - - if refund_token is None: + try: + # Token creation may swap proofs, so it is unsafe to retry. + refund_token = await send_token(amount, unit=unit, mint_url=mint) + except Exception as error: + logger.error( + "Failed to create refund token", + extra={ + "error": str(error), + "error_type": type(error).__name__, + "amount": amount, + "unit": unit, + "mint": mint, + }, + ) raise HTTPException( status_code=401, detail={ "error": { - "message": f"failed to create refund after {max_retries} attempts: {str(last_exception)}", + "message": f"failed to create refund: {error}", "type": "invalid_request_error", "code": "send_token_failed", } }, - ) + ) from error logger.info( "Refund token created successfully", @@ -3683,7 +3723,6 @@ class BaseUpstreamProvider: "amount": amount, "unit": unit, "mint": mint, - "attempt": attempt + 1, "token_preview": refund_token[:20] + "..." if len(refund_token) > 20 else refund_token, @@ -5362,7 +5401,7 @@ class BaseUpstreamProvider: except Exception as e: logger.error( f"Failed to refresh models cache for {self.provider_type or self.base_url}", - extra={"error": str(e), "error_type": type(e).__name__}, + extra={"error": repr(e), "error_type": type(e).__name__}, ) def get_cached_models(self) -> list[Model]: diff --git a/routstr/upstream/ppqai.py b/routstr/upstream/ppqai.py index 50ab4532..65ccbb57 100644 --- a/routstr/upstream/ppqai.py +++ b/routstr/upstream/ppqai.py @@ -1,5 +1,9 @@ from __future__ import annotations +import asyncio +import random +import time +from dataclasses import dataclass, field from typing import TYPE_CHECKING, Optional import httpx @@ -15,6 +19,87 @@ if TYPE_CHECKING: logger = get_logger(__name__) +_PPQ_SAFE_READ_ATTEMPTS = 3 +_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0 + + +class PPQCircuitOpenError(RuntimeError): + pass + + +@dataclass +class _PPQCircuitState: + consecutive_failures: int = 0 + cooldown_until: float = 0.0 + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + loop: asyncio.AbstractEventLoop | None = None + + +_ppq_circuits: dict[str, _PPQCircuitState] = {} + + +def _ppq_origin(url: str) -> str: + parsed = httpx.URL(url) + port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0) + return f"{parsed.scheme}://{parsed.host}:{port}" + + +async def _safe_read_request( + client: httpx.AsyncClient, + method: str, + url: str, + *, + headers: dict[str, str], + json: dict[str, object] | None = None, +) -> httpx.Response: + state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState()) + loop = asyncio.get_running_loop() + if state.loop is not loop: + # Locks cannot be reused across event loops. + state.lock = asyncio.Lock() + state.loop = loop + async with state.lock: + remaining = state.cooldown_until - time.monotonic() + if remaining > 0: + raise PPQCircuitOpenError( + f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s" + ) + + for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1): + try: + response = await client.request(method, url, headers=headers, json=json) + response.raise_for_status() + state.consecutive_failures = 0 + state.cooldown_until = 0.0 + return response + except (httpx.TransportError, httpx.HTTPStatusError) as error: + retryable_status = isinstance(error, httpx.HTTPStatusError) and ( + error.response.status_code in {502, 503, 504} + ) + if not isinstance(error, httpx.TransportError) and not retryable_status: + raise + state.consecutive_failures += 1 + if attempt >= _PPQ_SAFE_READ_ATTEMPTS: + state.cooldown_until = ( + time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS + ) + raise + base_delay = 0.25 * (2 ** (attempt - 1)) + delay = base_delay + random.uniform(0.0, base_delay) + logger.warning( + "PPQ.AI safe read failed; retrying", + extra={ + "url": url, + "attempt": attempt, + "max_attempts": _PPQ_SAFE_READ_ATTEMPTS, + "backoff_seconds": round(delay, 3), + "error": repr(error), + "error_type": type(error).__name__, + }, + ) + await asyncio.sleep(delay) + raise RuntimeError("unreachable") + class PPQAIModelPricing(BaseModel): ui: Optional[dict[str, float]] = None @@ -123,122 +208,109 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): url = f"{self.base_url}/models" headers = {"Authorization": f"Bearer {self.api_key}"} - try: - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - data = response.json() + async with httpx.AsyncClient(timeout=30.0) as client: + response = await _safe_read_request(client, "GET", url, headers=headers) + data = response.json() - models_data = data.get("data", []) + models_data = data.get("data", []) - or_models = [ - Model(**model) # type: ignore - for model in await async_fetch_openrouter_models() - ] + or_models = [ + Model(**model) # type: ignore + for model in await async_fetch_openrouter_models() + ] - models = [] - for model_data in models_data: - try: - ppqai_model = PPQAIModel.parse_obj(model_data) - if ppqai_model.id in self.IGNORED_MODEL_IDS: - continue + models = [] + for model_data in models_data: + try: + ppqai_model = PPQAIModel.parse_obj(model_data) + if ppqai_model.id in self.IGNORED_MODEL_IDS: + continue - or_model = next( - ( - model - for model in or_models - if (model.id == ppqai_model.id) - or (model.id.split("/")[-1] == ppqai_model.id) - or (model.id == ppqai_model.id.split("/")[-1]) - ), - None, - ) + or_model = next( + ( + model + for model in or_models + if (model.id == ppqai_model.id) + or (model.id.split("/")[-1] == ppqai_model.id) + or (model.id == ppqai_model.id.split("/")[-1]) + ), + None, + ) - if or_model: - input_price = None - if ppqai_model.pricing.api: - input_price = ppqai_model.pricing.api.get( - "input_per_1M" - ) - elif ppqai_model.pricing.input_per_1M_tokens: - input_price = ppqai_model.pricing.input_per_1M_tokens + if or_model: + input_price = None + if ppqai_model.pricing.api: + input_price = ppqai_model.pricing.api.get("input_per_1M") + elif ppqai_model.pricing.input_per_1M_tokens: + input_price = ppqai_model.pricing.input_per_1M_tokens - if input_price is not None: - or_model.pricing.prompt = input_price / 1_000_000 + if input_price is not None: + or_model.pricing.prompt = input_price / 1_000_000 - output_price = None - if ppqai_model.pricing.api: - output_price = ppqai_model.pricing.api.get( - "output_per_1M" - ) - elif ppqai_model.pricing.output_per_1M_tokens: - output_price = ppqai_model.pricing.output_per_1M_tokens + output_price = None + if ppqai_model.pricing.api: + output_price = ppqai_model.pricing.api.get("output_per_1M") + elif ppqai_model.pricing.output_per_1M_tokens: + output_price = ppqai_model.pricing.output_per_1M_tokens - if output_price is not None: - or_model.pricing.completion = output_price / 1_000_000 + if output_price is not None: + or_model.pricing.completion = output_price / 1_000_000 - if cl := ppqai_model.context_length: - or_model.context_length = cl - models.append(or_model) - else: - input_price = 0.0 - if ppqai_model.pricing.api: - input_price = ppqai_model.pricing.api.get( - "input_per_1M", 0.0 - ) - elif ppqai_model.pricing.input_per_1M_tokens: - input_price = ppqai_model.pricing.input_per_1M_tokens - - output_price = 0.0 - if ppqai_model.pricing.api: - output_price = ppqai_model.pricing.api.get( - "output_per_1M", 0.0 - ) - elif ppqai_model.pricing.output_per_1M_tokens: - output_price = ppqai_model.pricing.output_per_1M_tokens - - models.append( - Model( - id=ppqai_model.id, - name=ppqai_model.name, - created=ppqai_model.created_at // 1000, - description=f"{ppqai_model.provider or 'PPQ.AI'} model", - context_length=ppqai_model.context_length, - architecture=Architecture( - modality="text->text", - input_modalities=["text"], - output_modalities=["text"], - tokenizer="Unknown", - instruct_type=None, - ), - pricing=Pricing( - prompt=input_price / 1_000_000, - completion=output_price / 1_000_000, - request=0.0, - image=0.0, - web_search=0.0, - internal_reasoning=0.0, - ), - ) + if cl := ppqai_model.context_length: + or_model.context_length = cl + models.append(or_model) + else: + input_price = 0.0 + if ppqai_model.pricing.api: + input_price = ppqai_model.pricing.api.get( + "input_per_1M", 0.0 + ) + elif ppqai_model.pricing.input_per_1M_tokens: + input_price = ppqai_model.pricing.input_per_1M_tokens + + output_price = 0.0 + if ppqai_model.pricing.api: + output_price = ppqai_model.pricing.api.get( + "output_per_1M", 0.0 + ) + elif ppqai_model.pricing.output_per_1M_tokens: + output_price = ppqai_model.pricing.output_per_1M_tokens + + models.append( + Model( + id=ppqai_model.id, + name=ppqai_model.name, + created=ppqai_model.created_at // 1000, + description=f"{ppqai_model.provider or 'PPQ.AI'} model", + context_length=ppqai_model.context_length, + architecture=Architecture( + modality="text->text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="Unknown", + instruct_type=None, + ), + pricing=Pricing( + prompt=input_price / 1_000_000, + completion=output_price / 1_000_000, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + ), ) - except Exception as e: - logger.warning( - "Failed to parse PPQ.AI model", - extra={ - "model_id": model_data.get("id", "unknown"), - "error": str(e), - "error_type": type(e).__name__, - }, ) + except Exception as e: + logger.warning( + "Failed to parse PPQ.AI model", + extra={ + "model_id": model_data.get("id", "unknown"), + "error": str(e), + "error_type": type(e).__name__, + }, + ) - return models - - except Exception as e: - logger.error( - "Error fetching models from PPQ.AI", - extra={"error": str(e), "error_type": type(e).__name__}, - ) - return [] + return models async def on_upstream_error_redirect( self, status_code: int, error_message: str @@ -360,8 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): ) async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() + response = await _safe_read_request(client, "GET", url, headers=headers) status_data = response.json() is_paid = status_data.get("status") == "Settled" @@ -460,8 +531,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): logger.debug("Checking PPQ.AI account balance", extra={"url": url}) async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post(url, headers=headers, json={}) - response.raise_for_status() + response = await _safe_read_request( + client, "POST", url, headers=headers, json={} + ) balance_data = response.json() logger.debug( diff --git a/routstr/wallet.py b/routstr/wallet.py index c6cb1238..0c0ecc7c 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -14,6 +14,7 @@ from typing import AsyncGenerator, TypedDict import httpx from cashu.core.base import MeltQuote, MeltQuoteState, MintQuote, Proof, Token from cashu.core.mint_info import MintInfo as _CashuMintInfo +from cashu.wallet.crud import get_keysets as get_cashu_keysets from cashu.wallet.helpers import deserialize_token_from_string from cashu.wallet.wallet import Wallet as _CashuWallet from pydantic_core import PydanticUndefined @@ -121,6 +122,12 @@ def _mints_to_inspect() -> list[str]: return mint_urls +_WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS = 30 +_WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS = 300 +_mint_metadata_last_load: dict[str, float] = {} +_mint_metadata_load_locks: dict[str, asyncio.Lock] = {} + + class Wallet(_CashuWallet): """Cashu adapter that preserves HTTP 429 for Routstr's mint policy.""" @@ -141,11 +148,34 @@ class Wallet(_CashuWallet): _CashuWallet.raise_on_error_request(resp) async def load_mint( - self, keyset_id: str = "", force_old_keysets: bool = False + self, + keyset_id: str = "", + force_old_keysets: bool = False, + *, + force_refresh: bool = False, ) -> None: - await self.load_mint_keysets(force_old_keysets) - await self.activate_keyset(keyset_id) - await self.load_mint_info(reload=True) + mint_url = str(self.url) + lock = _mint_metadata_load_locks.setdefault(mint_url, asyncio.Lock()) + async with lock: + now = time.monotonic() + last = _mint_metadata_last_load.get(mint_url) + if ( + not force_refresh + and last is not None + and now - last < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS + ): + try: + await self.load_keysets_from_db() + await self.activate_keyset(keyset_id) + await self.load_mint_info(reload=False) + return + except Exception: + pass + + await self.load_mint_keysets(force_old_keysets) + await self.activate_keyset(keyset_id) + await self.load_mint_info(reload=True) + _mint_metadata_last_load[mint_url] = time.monotonic() class MintConnectionError(Exception): @@ -491,7 +521,11 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int async def _send_locked( - amount: int, unit: str, mint_url: str | None = None + amount: int, + unit: str, + mint_url: str | None = None, + *, + owner_only: bool = False, ) -> tuple[int, str]: effective_mint_url = await find_trusted_mint_with_funds( amount, unit, mint_url, force_reload=True @@ -501,6 +535,12 @@ async def _send_locked( wallet, effective_mint_url, unit, not_reserved=True ) proofs_for_mint = sum(proof.amount for proof in proofs) + if owner_only: + owner_balance = await _owner_balance_for_mint_and_unit( + effective_mint_url, unit, proofs_for_mint + ) + if owner_balance < amount: + raise ValueError("Owner Cashu balance is insufficient for auto top-up") all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit) reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved) @@ -547,6 +587,13 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str return token +async def send_token_from_owner_locked( + amount: int, unit: str, mint_url: str | None = None +) -> str: + _, token = await _send_locked(amount, unit, mint_url, owner_only=True) + return token + + class Bolt11PaymentNotAttempted(Exception): """The invoice was definitively not paid, so the attempt can be retried. @@ -1048,6 +1095,7 @@ async def _request_mint_with_fallback( mint_url, settings.primary_mint_unit, retry_on_rate_limit=False, + load_proofs=False, ) quote = await run_mint_operation( lambda: wallet.request_mint(amount), @@ -1179,6 +1227,7 @@ async def _calculate_swap_amount( lambda: token_wallet.melt_quote(dummy_mint_quote.request), op_name="swap_fee_est_melt_quote", mint_url=token_mint_url, + retry_timeouts=False, ) fee_reserve = dummy_melt_quote.fee_reserve @@ -1416,6 +1465,7 @@ async def swap_to_trusted_mint( lambda: token_wallet.melt_quote(mint_quote.request), op_name="swap_melt_quote", mint_url=token_obj.mint, + retry_timeouts=False, ) except Exception as error: if is_mint_connection_error(error): @@ -1788,21 +1838,34 @@ async def _credit_balance_locked( ) return amount except Exception as e: - logger.error( - "credit_balance: Error during token redemption", - extra={"error": str(e), "error_type": type(e).__name__}, + classification = classify_redemption_error(e) + expected_codes = { + "cashu_token_already_spent", + "cashu_source_mint_unreachable", + "cashu_mint_unreachable", + "cashu_mint_rate_limited", + } + log = ( + logger.info + if classification is not None and classification[3] in expected_codes + else logger.error + ) + log( + "credit_balance: Token redemption failed", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "error_code": classification[3] if classification else None, + }, ) raise _wallets: dict[str, Wallet] = {} +# Proofs require a shorter refresh interval than remote mint metadata. _wallet_last_load: dict[str, float] = {} +_wallet_last_mint_load: dict[str, float] = {} _wallet_load_locks: dict[str, asyncio.Lock] = {} -# Minimum seconds between full mint info + proof reloads for the same -# wallet. Prevents redundant mint API calls when get_wallet(load=True) -# is called rapidly by multiple background tasks (balance fetch, payout, -# auto-topup all hitting get_wallet within the same cycle). -_WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30 async def get_wallet( @@ -1811,8 +1874,9 @@ async def get_wallet( load: bool = True, retry_on_rate_limit: bool = True, force_reload: bool = False, + load_proofs: bool = True, ) -> Wallet: - global _wallets, _wallet_last_load, _wallet_load_locks + global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks id = f"{mint_url}_{unit}" lock = _wallet_load_locks.setdefault(id, asyncio.Lock()) async with lock: @@ -1821,25 +1885,39 @@ async def get_wallet( if load: now = time.monotonic() - last = _wallet_last_load.get(id) + last_mint_load = _wallet_last_mint_load.get(id) if ( force_reload - or last is None - or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS + or last_mint_load is None + or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS ): await run_mint_operation( - lambda: _wallets[id].load_mint(), + lambda: ( + _wallets[id].load_mint(force_refresh=True) + if force_reload + else _wallets[id].load_mint() + ), op_name="load_mint", mint_url=mint_url, retry_on_rate_limit=retry_on_rate_limit, ) - await run_mint_operation( - lambda: _wallets[id].load_proofs(reload=True), - op_name="load_proofs", - mint_url=mint_url, - retry_on_rate_limit=retry_on_rate_limit, - ) - _wallet_last_load[id] = time.monotonic() + _wallet_last_mint_load[id] = time.monotonic() + + if load_proofs: + last_proof_load = _wallet_last_load.get(id) + if ( + force_reload + or last_proof_load is None + or now - last_proof_load + >= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS + ): + await run_mint_operation( + lambda: _wallets[id].load_proofs(reload=True), + op_name="load_proofs", + mint_url=mint_url, + retry_on_rate_limit=retry_on_rate_limit, + ) + _wallet_last_load[id] = time.monotonic() return _wallets[id] @@ -1912,13 +1990,14 @@ async def _get_supported_mint_units(mint_url: str) -> list[str]: if cached is not None and now < cached[0]: return cached[1] - wallet = await get_wallet(mint_url, settings.primary_mint_unit, load=False) - keysets = await run_mint_operation( - lambda: wallet._get_keysets(), - op_name="get_mint_keysets", - mint_url=mint_url, + # A metadata load populates Cashu's shared keyset cache for all units. + wallet = await get_wallet( + mint_url, + settings.primary_mint_unit, retry_on_rate_limit=False, + load_proofs=False, ) + keysets = await get_cashu_keysets(mint_url=wallet.url, db=wallet.db) units: list[str] = [] for keyset in keysets: if not keyset.active or keyset.unit is None: diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 265c73c9..f369182c 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -1,7 +1,7 @@ import asyncio import json import os -from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple +from typing import Any, AsyncGenerator, Callable, Dict, Iterator, List, Optional, Tuple from unittest.mock import MagicMock, patch import pytest @@ -68,6 +68,14 @@ os.environ.pop("ADMIN_PASSWORD", None) from routstr.core.db import ApiKey, get_session # noqa: E402 from routstr.core.main import app, lifespan # noqa: E402 +from routstr.mint import MintRateGuard # noqa: E402 + + +@pytest.fixture(autouse=True) +def isolate_mint_rate_guards() -> Iterator[None]: + MintRateGuard._guards.clear() + yield + MintRateGuard._guards.clear() @pytest.fixture(scope="session") diff --git a/tests/integration/test_ppq_auto_topup_claim.py b/tests/integration/test_ppq_auto_topup_claim.py index 433389cd..05b99c03 100644 --- a/tests/integration/test_ppq_auto_topup_claim.py +++ b/tests/integration/test_ppq_auto_topup_claim.py @@ -23,6 +23,7 @@ from routstr.upstream.auto_topup import ( _ppq_request_id, _ppq_spent_last_24h_usd, _ppq_state_id_for_provider, + _reconcile_ppq_state, _record_ppq_invoice, _set_ppq_state_terminal, get_ppq_auto_topup_state, @@ -35,6 +36,8 @@ pytestmark = pytest.mark.asyncio def _row(provider_id: int = 1) -> MagicMock: row = MagicMock() row.id = provider_id + row.api_key = "secret" + row.provider_settings = None return row @@ -111,12 +114,35 @@ async def test_claim_is_reusable_once_the_previous_attempt_finished( await _seed_provider() first = await _claim_ppq_topup(_row()) assert first is not None - assert await _set_ppq_state_terminal(_row(), first, collected=True, swept=False) + assert await _set_ppq_state_terminal(_row(), first, collected=False, swept=True) second = await _claim_ppq_topup(_row()) assert second is not None and second != first +async def test_settled_claim_suppresses_immediate_duplicate( + patched_db_engine: Any, +) -> None: + await _seed_provider() + row = _row() + operation_id = await _claim_ppq_topup(row) + assert operation_id is not None + assert await _set_ppq_state_terminal(row, operation_id, collected=True, swept=False) + + assert await _reconcile_ppq_state(row, provider=None) is True + assert await _claim_ppq_topup(row) is None + + +async def test_claim_rejects_stale_provider_configuration( + patched_db_engine: Any, +) -> None: + await _seed_provider() + stale = _row() + stale.provider_settings = '{"auto_topup":true}' + + assert await _claim_ppq_topup(stale) is None + + async def test_recording_the_invoice_moves_the_claim_in_flight( patched_db_engine: Any, ) -> None: @@ -287,7 +313,13 @@ async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim( assert audit["collected"] is True assert "lnbc-secret-invoice" not in audit["token"] - # Reusing the deterministic claim lock must not overwrite history. + assert await _claim_ppq_topup(_row()) is None + async with create_session() as session: + state = await session.get(CashuTransaction, _ppq_state_id_for_provider(1)) + assert state is not None + state.created_at = int(time.time()) - 301 + session.add(state) + await session.commit() assert await _claim_ppq_topup(_row()) is not None async with create_session() as session: assert await session.get(CashuTransaction, audit["id"]) is not None diff --git a/tests/integration/test_provider_management.py b/tests/integration/test_provider_management.py index b7db0c6c..de9a3112 100644 --- a/tests/integration/test_provider_management.py +++ b/tests/integration/test_provider_management.py @@ -686,7 +686,7 @@ async def test_no_database_changes_during_provider_operations( @pytest.mark.integration @pytest.mark.asyncio -async def test_admin_routstr_topup_retries_transient_upstream_failure( +async def test_admin_routstr_topup_does_not_duplicate_invoice_on_upstream_failure( integration_client: AsyncClient, integration_session: Any, ) -> None: @@ -739,16 +739,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure( assert json["api_key"] == "sk-upstream-test" assert headers["Authorization"] == "Bearer sk-upstream-test" - if self.calls == 1: - return MockResponse(500, {"detail": "warmup failure"}) - - return MockResponse( - 200, - { - "bolt11": "lnbc1testinvoice", - "invoice_id": "invoice-123", - }, - ) + return MockResponse(500, {"detail": "ambiguous upstream failure"}) mock_client = MockAsyncClient() @@ -759,11 +750,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure( json={"amount": 10}, ) - assert response.status_code == 200 - data = response.json() - assert data["ok"] is True - assert data["topup_data"]["payment_request"] == "lnbc1testinvoice" - assert data["topup_data"]["invoice_id"] == "invoice-123" - assert mock_client.calls == 2 + assert response.status_code == 500 + assert mock_client.calls == 1 finally: admin_sessions.pop(admin_token, None) diff --git a/tests/integration/test_routstr_auto_topup_claim.py b/tests/integration/test_routstr_auto_topup_claim.py index 25b827a9..dc246a63 100644 --- a/tests/integration/test_routstr_auto_topup_claim.py +++ b/tests/integration/test_routstr_auto_topup_claim.py @@ -25,6 +25,7 @@ from routstr.upstream.auto_topup import ( _check_and_topup, _claim_routstr_topup, _parse_routstr_request_id, + _persist_routstr_token_and_mark_sent, _routstr_spent_last_24h_sats, _routstr_state_id_for_provider, get_routstr_auto_topup_state, @@ -75,7 +76,9 @@ def _patch_wallet(module: Any, peer: MagicMock, token: str) -> ExitStack: patch.object(module.RoutstrUpstreamProvider, "from_db_row", return_value=peer) ) stack.enter_context( - patch.object(module, "send_token", AsyncMock(return_value=token)) + patch.object( + module, "send_token_from_owner_locked", AsyncMock(return_value=token) + ) ) stack.enter_context( patch.object(module, "token_mint_url", return_value="https://mint.test") @@ -111,6 +114,32 @@ async def test_second_worker_cannot_claim_while_the_first_holds_one( assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is None +async def test_token_and_sent_claim_roll_back_together_on_commit_failure( + patched_db_engine: Any, +) -> None: + row = await _seed_provider() + operation_id = await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) + assert operation_id is not None + + with patch( + "sqlmodel.ext.asyncio.session.AsyncSession.commit", + new=AsyncMock(side_effect=RuntimeError("commit failed")), + ): + with pytest.raises(RuntimeError, match="commit failed"): + await _persist_routstr_token_and_mark_sent( + row, + operation_id, + expected_sats=TOPUP_SATS, + token="cashu-token-atomic", + amount=TOPUP_SATS, + mint_url="https://mint.test", + ) + + claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr] + assert claim is not None and claim.phase != ROUTSTR_PHASE_SENT + assert await _sent_tokens() == [] + + async def test_token_is_persisted_before_it_reaches_the_peer( patched_db_engine: Any, ) -> None: @@ -142,7 +171,7 @@ async def test_untracked_token_is_returned_and_never_sent( _patch_wallet(auto_topup_module, peer, "cashu-token-1"), patch.object( auto_topup_module, - "store_cashu_transaction", + "_persist_routstr_token_and_mark_sent", AsyncMock(side_effect=RuntimeError("database unavailable")), ), patch.object( diff --git a/tests/integration/test_wallet_authentication.py b/tests/integration/test_wallet_authentication.py index 9dc19947..26407701 100644 --- a/tests/integration/test_wallet_authentication.py +++ b/tests/integration/test_wallet_authentication.py @@ -3,6 +3,7 @@ Integration tests for wallet authentication system including API key generation Tests POST /v1/wallet/topup endpoint and authorization header validation. """ +import asyncio from datetime import datetime, timedelta from typing import Any @@ -113,38 +114,32 @@ async def test_api_key_generation_invalid_token( async def test_duplicate_token_handling( integration_client: AsyncClient, testmint_wallet: Any, db_snapshot: Any ) -> None: - """Test that duplicate tokens return the same API key without double-spending""" - - # Generate a valid token - amount = 500 # 500 sats + amount = 500 token = await testmint_wallet.mint_tokens(amount) - - # First use of token integration_client.headers["Authorization"] = f"Bearer {token}" - response1 = await integration_client.get("/v1/wallet/info") - assert response1.status_code == 200 + + response1, response2 = await asyncio.gather( + integration_client.get("/v1/wallet/info"), + integration_client.get("/v1/wallet/info"), + ) + assert response1.status_code < 500 + assert response2.status_code < 500 + assert response1.status_code == response2.status_code == 200 api_key1 = response1.json()["api_key"] - balance1 = response1.json()["balance"] - - # Capture state after first submission - await db_snapshot.capture() - - # Second use of same token - should return same API key since it's already created - response2 = await integration_client.get("/v1/wallet/info") - assert response2.status_code == 200 api_key2 = response2.json()["api_key"] + balance1 = response1.json()["balance"] balance2 = response2.json()["balance"] - - # Should return the same API key and balance assert api_key1 == api_key2 - assert balance1 == balance2 + assert balance1 == balance2 == amount * 1000 - # Verify no additional database changes + await db_snapshot.capture() + replay = await integration_client.get("/v1/wallet/info") + assert replay.status_code == 200 + assert replay.json()["api_key"] == api_key1 diff = await db_snapshot.diff() assert len(diff["api_keys"]["added"]) == 0 assert len(diff["api_keys"]["modified"]) == 0 - # Original API key should still work with original balance integration_client.headers["Authorization"] = f"Bearer {api_key1}" wallet_response = await integration_client.get("/v1/wallet/") assert wallet_response.status_code == 200 diff --git a/tests/integration/test_wallet_melt_restart.py b/tests/integration/test_wallet_melt_restart.py index 5bb607fa..9d75f542 100644 --- a/tests/integration/test_wallet_melt_restart.py +++ b/tests/integration/test_wallet_melt_restart.py @@ -10,9 +10,12 @@ invalidate them on "paid" or release them on "unpaid". """ from pathlib import Path +from unittest.mock import AsyncMock, patch +import httpx import pytest -from cashu.core.base import Proof +from cashu.core.base import MeltQuote, MeltQuoteState, Proof +from cashu.core.models import PostMeltQuoteResponse from cashu.wallet import crud from cashu.wallet.wallet import Wallet @@ -47,6 +50,20 @@ async def _seed_ambiguous_melt(wallet: Wallet) -> list[Proof]: proofs = [_proof("secret-a"), _proof("secret-b", amount=32)] for proof in proofs: await crud.store_proof(proof, db=wallet.db) + await crud.store_bolt11_melt_quote( + db=wallet.db, + quote=MeltQuote( + quote=QUOTE_ID, + method="bolt11", + request="lnbc1-test", + checking_id="", + unit="sat", + amount=95, + fee_reserve=1, + state=MeltQuoteState.pending, + mint=str(wallet.url), + ), + ) await wallet.set_reserved_for_melt(proofs, reserved=True, quote_id=QUOTE_ID) # cashu's `except` block in melt(): @@ -75,6 +92,34 @@ async def test_melt_recovery_is_findable_by_quote_after_restart( assert all(p.melt_id == QUOTE_ID for p in found) +async def test_paid_reconciliation_invalidates_recovered_proofs_after_restart( + tmp_path: Path, +) -> None: + wallet = await _wallet(tmp_path) + await _seed_ambiguous_melt(wallet) + + restarted = await _wallet(tmp_path) + remote = PostMeltQuoteResponse( + quote=QUOTE_ID, + amount=95, + unit="sat", + request="lnbc1-test", + fee_reserve=1, + state=MeltQuoteState.paid.value, + expiry=None, + payment_preimage="preimage", + ) + with patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(return_value=remote), + ): + reconciled = await restarted.get_melt_quote(QUOTE_ID) + + assert reconciled is not None and reconciled.state == MeltQuoteState.paid + assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == [] + assert await crud.get_proofs(db=restarted.db) == [] + + async def test_send_style_reservation_would_not_be_reconcilable( tmp_path: Path, ) -> None: @@ -97,19 +142,64 @@ async def test_send_style_reservation_would_not_be_reconcilable( async def test_unpaid_reconciliation_releases_recovered_proofs_after_restart( tmp_path: Path, ) -> None: - """The full recovery arc: crash, restart, mint says unpaid, funds usable.""" wallet = await _wallet(tmp_path) await _seed_ambiguous_melt(wallet) restarted = await _wallet(tmp_path) - found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) - assert len(found) == 2 + remote = PostMeltQuoteResponse( + quote=QUOTE_ID, + amount=95, + unit="sat", + request="lnbc1-test", + fee_reserve=1, + state=MeltQuoteState.unpaid.value, + expiry=None, + ) + with patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(return_value=remote), + ): + reconciled = await restarted.get_melt_quote(QUOTE_ID) - # What get_melt_quote() does on an "unpaid" answer. - await restarted.set_reserved_for_melt(found, reserved=False, quote_id=None) - - released = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) - assert released == [] + assert reconciled is not None and reconciled.state == MeltQuoteState.unpaid + assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == [] all_proofs = await crud.get_proofs(db=restarted.db) assert len(all_proofs) == 2 - assert all(not p.reserved for p in all_proofs) # spendable again + assert all(not p.reserved for p in all_proofs) + + +async def test_pending_and_transport_reconciliation_keep_recovered_reservation( + tmp_path: Path, +) -> None: + wallet = await _wallet(tmp_path) + await _seed_ambiguous_melt(wallet) + restarted = await _wallet(tmp_path) + pending = PostMeltQuoteResponse( + quote=QUOTE_ID, + amount=95, + unit="sat", + request="lnbc1-test", + fee_reserve=1, + state=MeltQuoteState.pending.value, + expiry=None, + ) + + with patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(return_value=pending), + ): + reconciled = await restarted.get_melt_quote(QUOTE_ID) + assert reconciled is not None and reconciled.state == MeltQuoteState.pending + found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) + assert len(found) == 2 and all(proof.reserved for proof in found) + + with ( + patch( + "cashu.wallet.v1_api.LedgerAPI.get_melt_quote", + new=AsyncMock(side_effect=httpx.ReadTimeout("mint unavailable")), + ), + pytest.raises(httpx.ReadTimeout), + ): + await restarted.get_melt_quote(QUOTE_ID) + found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) + assert len(found) == 2 and all(proof.reserved for proof in found) diff --git a/tests/unit/test_auto_topup.py b/tests/unit/test_auto_topup.py index d28f939c..349dc9ef 100644 --- a/tests/unit/test_auto_topup.py +++ b/tests/unit/test_auto_topup.py @@ -1,4 +1,6 @@ import json +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -65,7 +67,9 @@ async def test_auto_topup_refuses_invalid_settings_before_touching_the_wallet() "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", return_value=provider, ), - patch("routstr.upstream.auto_topup.send_token", AsyncMock()) as send, + patch( + "routstr.upstream.auto_topup.send_token_from_owner_locked", AsyncMock() + ) as send, ): await _check_and_topup(row) @@ -73,6 +77,63 @@ async def test_auto_topup_refuses_invalid_settings_before_touching_the_wallet() send.assert_not_awaited() +@pytest.mark.asyncio +async def test_routstr_outgoing_audit_is_persisted_under_wallet_guard() -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(return_value=0.0) + provider.topup = AsyncMock(return_value={"error": "stop after send"}) + inside_guard = False + + @asynccontextmanager + async def guard() -> AsyncIterator[None]: + nonlocal inside_guard + inside_guard = True + try: + yield + finally: + inside_guard = False + + async def send(*_args: object) -> str: + assert inside_guard + return "cashu-token" + + async def persist(*_args: object, **_kwargs: object) -> None: + assert inside_guard + + with ( + patch( + "routstr.upstream.auto_topup.RoutstrUpstreamProvider.from_db_row", + return_value=provider, + ), + patch( + "routstr.upstream.auto_topup._reconcile_routstr_state", + AsyncMock(return_value=False), + ), + patch( + "routstr.upstream.auto_topup._routstr_spent_last_24h_sats", + AsyncMock(return_value=0), + ), + patch( + "routstr.upstream.auto_topup._claim_routstr_topup", + AsyncMock(return_value="operation-1"), + ), + patch("routstr.upstream.auto_topup.wallet_operation_guard", side_effect=guard), + patch( + "routstr.upstream.auto_topup.send_token_from_owner_locked", + side_effect=send, + ), + patch( + "routstr.upstream.auto_topup._persist_routstr_token_and_mark_sent", + side_effect=persist, + ), + patch( + "routstr.upstream.auto_topup.token_mint_url", + return_value="https://mint.test", + ), + ): + await _check_and_topup(_row()) + + def _ppq_row() -> MagicMock: row = MagicMock() row.id = "ppq-provider-1" @@ -443,6 +504,30 @@ async def test_ppq_auto_topup_skips_when_balance_meets_threshold() -> None: provider.initiate_topup.assert_not_awaited() +@pytest.mark.asyncio +async def test_ppq_auto_topup_requires_two_below_threshold_reads() -> None: + provider = MagicMock() + provider.get_balance = AsyncMock(side_effect=[2.5, 5.0]) + provider.initiate_topup = AsyncMock() + + with ( + patch( + "routstr.upstream.auto_topup.PPQAIUpstreamProvider.from_db_row", + return_value=provider, + ), + patch( + "routstr.upstream.auto_topup._reconcile_ppq_state", + AsyncMock(return_value=False), + ), + patch("routstr.upstream.auto_topup._claim_ppq_topup", AsyncMock()) as claim, + ): + await _check_and_topup(_ppq_row()) + + assert provider.get_balance.await_count == 2 + claim.assert_not_awaited() + provider.initiate_topup.assert_not_awaited() + + @pytest.mark.asyncio async def test_ppq_auto_topup_skips_when_daily_spend_cap_reached() -> None: provider = MagicMock() diff --git a/tests/unit/test_core_exceptions.py b/tests/unit/test_core_exceptions.py index 3d9469df..5f4cc577 100644 --- a/tests/unit/test_core_exceptions.py +++ b/tests/unit/test_core_exceptions.py @@ -1,4 +1,5 @@ import json +from unittest.mock import patch import pytest from fastapi import HTTPException @@ -34,11 +35,14 @@ async def test_structured_http_error_uses_standard_error_envelope() -> None: "details": {"mint": "https://mint.example"}, } - response = await http_exception_handler( - request, - HTTPException(status_code=503, detail={"error": error}), - ) + with patch("routstr.core.exceptions.logger") as logger: + response = await http_exception_handler( + request, + HTTPException(status_code=503, detail={"error": error}), + ) + logger.warning.assert_called_once() + logger.error.assert_not_called() assert response.status_code == 503 assert json.loads(response.body) == { "detail": {"error": error}, diff --git a/tests/unit/test_db_pool_config.py b/tests/unit/test_db_pool_config.py index 8f5d0367..bc2a8dd9 100644 --- a/tests/unit/test_db_pool_config.py +++ b/tests/unit/test_db_pool_config.py @@ -56,9 +56,40 @@ def test_non_sqlite_backend_enables_pre_ping_automatically( assert created is fake_engine assert factory.call_args.kwargs["pool_pre_ping"] is True + assert "timeout" not in factory.call_args.kwargs["connect_args"] assert listen.call_count == 2 +def test_file_sqlite_sets_busy_timeout_connect_arg( + monkeypatch: pytest.MonkeyPatch, tmp_path: object +) -> None: + monkeypatch.setattr(settings, "database_busy_timeout", 42.0) + fake_engine = MagicMock() + + with ( + patch.object(db, "create_async_engine", return_value=fake_engine) as factory, + patch.object(db.event, "listen"), + ): + create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/busy.db") + + assert factory.call_args.kwargs["connect_args"]["timeout"] == 42.0 + + +def test_memory_sqlite_omits_busy_timeout_connect_arg( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(settings, "database_busy_timeout", 42.0) + fake_engine = MagicMock() + + with ( + patch.object(db, "create_async_engine", return_value=fake_engine) as factory, + patch.object(db.event, "listen"), + ): + create_db_engine("sqlite+aiosqlite://") + + assert "timeout" not in factory.call_args.kwargs["connect_args"] + + @pytest.mark.asyncio async def test_every_created_engine_warns_for_long_checkouts( monkeypatch: pytest.MonkeyPatch, tmp_path: object diff --git a/tests/unit/test_fetch_all_balances.py b/tests/unit/test_fetch_all_balances.py index 89813cac..12e57fd4 100644 --- a/tests/unit/test_fetch_all_balances.py +++ b/tests/unit/test_fetch_all_balances.py @@ -136,19 +136,20 @@ async def test_supported_mint_units_come_from_active_keysets() -> None: msat = MagicMock(active=False, unit="msat") usd = MagicMock(active=True) usd.unit.name = "usd" - wallet = MagicMock() - wallet._get_keysets = AsyncMock(return_value=[usd, msat, sat]) + wallet = MagicMock(url="http://mint:3338", db=MagicMock()) + get_keysets = AsyncMock(return_value=[usd, msat, sat]) with ( patch.object(settings, "primary_mint_unit", "sat"), patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.wallet.get_cashu_keysets", get_keysets), ): units = await _get_supported_mint_units("http://mint:3338") cached_units = await _get_supported_mint_units("http://mint:3338") assert units == ["sat", "usd"] assert cached_units == units - wallet._get_keysets.assert_awaited_once() + get_keysets.assert_awaited_once_with(mint_url=wallet.url, db=wallet.db) @pytest.mark.asyncio diff --git a/tests/unit/test_lightning_settlement.py b/tests/unit/test_lightning_settlement.py index dbec3bd7..eddfde57 100644 --- a/tests/unit/test_lightning_settlement.py +++ b/tests/unit/test_lightning_settlement.py @@ -1,6 +1,6 @@ import asyncio import time -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Iterator from contextlib import asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, Mock, patch @@ -20,9 +20,17 @@ from routstr.lightning import ( get_invoice_status, recover_invoice, ) +from routstr.mint import MintRateGuard from routstr.wallet import Wallet +@pytest.fixture(autouse=True) +def _clear_mint_rate_guards() -> Iterator[None]: + MintRateGuard._guards.clear() + yield + MintRateGuard._guards.clear() + + def _invoice(**overrides: object) -> SimpleNamespace: values = { "id": "invoice-1", diff --git a/tests/unit/test_lnurl_amount_and_destination.py b/tests/unit/test_lnurl_amount_and_destination.py index 5ba459d9..e722141a 100644 --- a/tests/unit/test_lnurl_amount_and_destination.py +++ b/tests/unit/test_lnurl_amount_and_destination.py @@ -5,6 +5,7 @@ hand an LNURL a set of unreserved proofs, and reserving only after that check is what stops a pre-dispatch failure from stranding proofs. """ +import math from collections.abc import Callable from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -27,19 +28,28 @@ LNURL_DATA = { "max_sendable": 100_000_000, } -# 1000 sat minus the 11 sat estimated fee reserve. -EXPECTED_QUOTE_SAT = 989 +EXPECTED_QUOTE_SAT = 999 def _wallet( - quote_amount: int = EXPECTED_QUOTE_SAT, + quote_amount: int | None = None, ) -> tuple[MagicMock, list[MagicMock]]: - proofs = [MagicMock(amount=1000)] + proofs = [MagicMock(amount=1000, reserved=False)] wallet = MagicMock(url="https://mint.test") - wallet.melt_quote = AsyncMock( - return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount) - ) + wallet.get_fees_for_proofs.return_value = 0 + if quote_amount is None: + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=1, quote="q", amount=1000), + MagicMock(fee_reserve=1, quote="q", amount=EXPECTED_QUOTE_SAT), + ] + ) + else: + wallet.melt_quote = AsyncMock( + return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount) + ) wallet.select_to_send = AsyncMock(return_value=(proofs, None)) + wallet.get_fees_for_proofs = MagicMock(return_value=0) wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid)) wallet.set_reserved_for_send = AsyncMock() return wallet, proofs @@ -124,16 +134,21 @@ async def test_raw_send_to_lnurl_accepts_exact_invoice() -> None: ) assert paid == EXPECTED_QUOTE_SAT * 1000 - wallet.select_to_send.assert_awaited_once() + wallet.select_to_send.assert_not_awaited() + wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=True) wallet.melt.assert_awaited_once() @pytest.mark.asyncio async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None: - # 1_000_000 msat minus an 11 sat fee reserve leaves 989_000 msat. - wallet, proofs = _wallet(989_000) + wallet, proofs = _wallet() proofs[0].amount = 1_000_000 - wallet.select_to_send = AsyncMock(return_value=(proofs, None)) + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=1, quote="q", amount=1_000_000), + MagicMock(fee_reserve=1, quote="q", amount=999_999), + ] + ) data_patch, invoice_patch = _lnurl_patches() with data_patch, invoice_patch: @@ -141,7 +156,46 @@ async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None: wallet, proofs, "owner@ln.tld", "msat", amount=1_000_000 ) - assert paid == 989_000 + assert paid == 999_999 + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion() -> ( + None +): + proofs = [MagicMock(amount=1, reserved=False) for _ in range(1500)] + wallet = MagicMock(url="https://mint.test") + wallet.get_fees_for_proofs = MagicMock( + side_effect=lambda selected: math.ceil(len(selected) / 100) + ) + wallet.melt_quote = AsyncMock( + side_effect=[ + MagicMock(fee_reserve=10, quote="q1", amount=1500), + MagicMock(fee_reserve=10, quote="q2", amount=1475), + ] + ) + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid)) + wallet.set_reserved_for_send = AsyncMock() + checkpoint = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with data_patch, invoice_patch: + paid = await raw_send_to_lnurl( + wallet, + proofs, + "owner@ln.tld", + "sat", + amount=1500, + on_melt_quote=checkpoint, + ) + + assert paid == 1_475_000 + assert wallet.melt_quote.await_count == 2 + checkpoint.assert_awaited_once_with("q2") + wallet.select_to_send.assert_not_called() + selected = wallet.melt.await_args.kwargs["proofs"] + assert sum(proof.amount for proof in selected) == 1500 + assert 1475 + 10 + wallet.get_fees_for_proofs(selected) == 1500 @pytest.mark.asyncio @@ -291,3 +345,23 @@ async def test_send_to_lnurl_does_not_reserve_before_lnurl_validation() -> None: assert raw_send.await_args is not None assert raw_send.await_args.args[1] is proofs assert raw_send.await_args.kwargs["amount"] == 1000 + + +def test_select_melt_proofs_stops_at_minimal_cover_when_over_budget() -> None: + from routstr.payment.lnurl import _select_melt_proofs + + wallet = MagicMock() + wallet.get_fees_for_proofs = MagicMock(side_effect=lambda selected: len(selected)) + proofs = [MagicMock(amount=600, reserved=False) for _ in range(3)] + + selected, shortfall = _select_melt_proofs( + wallet, + proofs, + quote_amount=1000, + fee_reserve=0, + gross_budget=1000, + ) + + assert selected is None + assert shortfall == 2 + assert wallet.get_fees_for_proofs.call_count == 2 diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index dc35ab4c..ede3f5af 100644 --- a/tests/unit/test_lnurl_melt_timeout.py +++ b/tests/unit/test_lnurl_melt_timeout.py @@ -1,6 +1,7 @@ """LNURL melt attempts must not misclassify ambiguous payment outcomes.""" import asyncio +from collections.abc import Iterator from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -11,10 +12,19 @@ from cashu.core.base import MeltQuoteState from routstr.core.settings import settings from routstr.mint import MintCooldownError, MintRateGuard from routstr.payment.lnurl import ( + LNURLError, MeltOutcomeAmbiguousError, raw_send_to_lnurl, ) + +@pytest.fixture(autouse=True) +def _clear_mint_guards() -> Iterator[None]: + MintRateGuard._guards.clear() + yield + MintRateGuard._guards.clear() + + LNURL_DATA = { "callback_url": "https://ln.tld/cb", "min_sendable": 1_000, @@ -22,17 +32,23 @@ LNURL_DATA = { } -# 1000 sat minus the 11 sat estimated fee reserve. -QUOTE_AMOUNT_SAT = 989 +QUOTE_AMOUNT_SAT = 999 def _wallet() -> tuple[MagicMock, list[MagicMock]]: - proofs = [MagicMock(amount=1000)] + proofs = [MagicMock(amount=1000, reserved=False)] wallet = MagicMock(url="https://mint.test") wallet.melt_quote = AsyncMock( - return_value=MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT) + side_effect=[ + MagicMock(fee_reserve=1, quote="q", amount=1000), + MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT), + ] ) + wallet.melt = AsyncMock() wallet.select_to_send = AsyncMock(return_value=(proofs, None)) + wallet.get_fees_for_proofs = MagicMock(return_value=0) + wallet.set_reserved_for_melt = AsyncMock() + wallet.set_reserved_for_send = AsyncMock() return wallet, proofs @@ -50,7 +66,26 @@ def _lnurl_patches() -> tuple[Any, Any]: @pytest.mark.asyncio -async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> None: +async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.unpaid)) + wallet.get_melt_quote = AsyncMock() + data_patch, invoice_patch = _lnurl_patches() + + with ( + data_patch, + invoice_patch, + pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised, + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + assert not isinstance(raised.value, MeltOutcomeAmbiguousError) + wallet.get_melt_quote.assert_not_awaited() + wallet.set_reserved_for_send.assert_any_await(proofs, reserved=False) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_timeout_then_unpaid_remains_ambiguous() -> None: wallet, proofs = _wallet() async def _hang(**kwargs: object) -> None: @@ -67,12 +102,65 @@ async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> Non patch.object(settings, "mint_retry_max_attempts", 0), data_patch, invoice_patch, - pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"), + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), ): await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) wallet.get_melt_quote.assert_awaited_once_with("q") - wallet.set_reserved_for_melt.assert_not_called() + assert wallet.set_reserved_for_melt.await_count == 2 + wallet.set_reserved_for_melt.assert_awaited_with( + proofs, reserved=True, quote_id="q" + ) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_wrapped_transport_unpaid_remains_ambiguous() -> None: + wallet, proofs = _wallet() + + async def _wrapped_transport_error(**kwargs: object) -> None: + try: + raise httpx.ReadTimeout("response lost") + except httpx.ReadTimeout as transport_error: + raise Exception("could not pay invoice") from transport_error + + wallet.melt = AsyncMock(side_effect=_wrapped_transport_error) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_retry_max_attempts", 3), + data_patch, + invoice_patch, + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.melt.assert_awaited_once() + wallet.get_melt_quote.assert_awaited_once_with("q") + assert wallet.set_reserved_for_melt.await_count == 2 + wallet.set_reserved_for_melt.assert_awaited_with( + proofs, reserved=True, quote_id="q" + ) + + +@pytest.mark.asyncio +async def test_raw_send_to_lnurl_does_not_retry_melt_quote_timeout() -> None: + wallet, proofs = _wallet() + wallet.melt_quote = AsyncMock(side_effect=httpx.ReadTimeout("response lost")) + data_patch, invoice_patch = _lnurl_patches() + + with ( + patch.object(settings, "mint_retry_max_attempts", 3), + data_patch, + invoice_patch, + pytest.raises(httpx.TimeoutException), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.melt_quote.assert_awaited_once() + wallet.melt.assert_not_awaited() @pytest.mark.asyncio @@ -98,6 +186,9 @@ async def test_raw_send_to_lnurl_timeout_reconciled_paid_is_success() -> None: assert paid > 0 wallet.get_melt_quote.assert_awaited_once_with("q") + wallet.set_reserved_for_melt.assert_awaited_once_with( + proofs, reserved=True, quote_id="q" + ) @pytest.mark.asyncio @@ -120,6 +211,45 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None: wallet.get_melt_quote.assert_awaited_once_with("q") +@pytest.mark.asyncio +async def test_pending_then_immediate_unpaid_remains_reserved_and_ambiguous() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending)) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + data_patch, invoice_patch = _lnurl_patches() + + with ( + data_patch, + invoice_patch, + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + wallet.set_reserved_for_melt.assert_awaited_once_with( + proofs, reserved=True, quote_id="q" + ) + + +@pytest.mark.asyncio +async def test_immediate_unpaid_reservation_failure_stays_ambiguous() -> None: + wallet, proofs = _wallet() + wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending)) + wallet.get_melt_quote = AsyncMock( + return_value=MagicMock(state=MeltQuoteState.unpaid) + ) + wallet.set_reserved_for_melt = AsyncMock(side_effect=OSError("db locked")) + data_patch, invoice_patch = _lnurl_patches() + + with ( + data_patch, + invoice_patch, + pytest.raises(MeltOutcomeAmbiguousError, match="could not be restored"), + ): + await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) + + @pytest.mark.asyncio @pytest.mark.parametrize("rate_error", ["cooldown", "http_429"]) async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs( @@ -153,7 +283,8 @@ async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs( 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) + assert wallet.set_reserved_for_send.await_count == 2 + wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False) @pytest.mark.asyncio @@ -178,7 +309,8 @@ async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None: 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) + assert wallet.set_reserved_for_send.await_count == 2 + wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False) MintRateGuard._guards.pop(str(wallet.url), None) diff --git a/tests/unit/test_mint.py b/tests/unit/test_mint.py index 6b70a558..f6873fa4 100644 --- a/tests/unit/test_mint.py +++ b/tests/unit/test_mint.py @@ -10,6 +10,7 @@ from routstr.mint import ( MintRateGuard, MintRateLimitedError, fail_fast_mint_operations, + run_mint_operation, ) from routstr.wallet import Wallet @@ -103,6 +104,82 @@ async def test_cashu_429_dispatches_through_wallet_override() -> None: await wallet.mint_quote(1, Unit.sat) +@pytest.mark.asyncio +async def test_wrapped_transport_failure_opens_central_cooldown() -> None: + mint_url = "https://transport-failure.test" + MintRateGuard._guards.pop(mint_url, None) + + async def wrapped_failure() -> None: + try: + raise httpx.ReadTimeout("body stalled") + except httpx.ReadTimeout as error: + raise Exception("wallet wrapper") from error + + with pytest.raises(Exception, match="wallet wrapper"): + await run_mint_operation( + wrapped_failure, + mint_url=mint_url, + retry_timeouts=False, + ) + + guard = MintRateGuard.get(mint_url) + assert guard.cooldown_remaining() > 29 + probe = AsyncMock() + async with fail_fast_mint_operations(): + with pytest.raises(MintCooldownError): + await guard.run(probe) + probe.assert_not_awaited() + MintRateGuard._guards.pop(mint_url, None) + + +@pytest.mark.asyncio +async def test_timeout_retry_succeeds_without_opening_cooldown() -> None: + from routstr.core.settings import settings + + mint_url = "https://retryable-timeout.test" + MintRateGuard._guards.pop(mint_url, None) + calls = 0 + + async def flaky() -> str: + nonlocal calls + calls += 1 + if calls == 1: + raise httpx.ReadTimeout("first attempt stalled") + return "ok" + + with ( + patch.object(settings, "mint_retry_max_attempts", 2), + patch("routstr.mint.asyncio.sleep", AsyncMock()), + ): + result = await run_mint_operation(flaky, mint_url=mint_url) + + assert result == "ok" + assert calls == 2 + assert MintRateGuard.get(mint_url).cooldown_remaining() == 0.0 + MintRateGuard._guards.pop(mint_url, None) + + +@pytest.mark.asyncio +async def test_exhausted_timeout_retries_open_transport_cooldown() -> None: + from routstr.core.settings import settings + + mint_url = "https://exhausted-timeout.test" + MintRateGuard._guards.pop(mint_url, None) + + async def always_timeout() -> None: + raise httpx.ReadTimeout("stalled") + + with ( + patch.object(settings, "mint_retry_max_attempts", 1), + patch("routstr.mint.asyncio.sleep", AsyncMock()), + pytest.raises(httpx.TimeoutException), + ): + await run_mint_operation(always_timeout, mint_url=mint_url) + + assert MintRateGuard.get(mint_url).cooldown_remaining() > 29 + MintRateGuard._guards.pop(mint_url, None) + + async def test_guard_concurrency_change_preserves_cooldown_state() -> None: from routstr.core.settings import settings diff --git a/tests/unit/test_mint_fallback_trust.py b/tests/unit/test_mint_fallback_trust.py index f7809458..5f2a2ab6 100644 --- a/tests/unit/test_mint_fallback_trust.py +++ b/tests/unit/test_mint_fallback_trust.py @@ -1,7 +1,8 @@ """Persisted mint preferences must not bypass the configured trusted set.""" -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from routstr.core.settings import settings @@ -31,6 +32,23 @@ async def test_untrusted_allowed_mints_fall_back_to_trusted_set() -> None: assert attempted == [TRUSTED] +async def test_mint_quote_timeout_is_not_retried() -> None: + wallet = MagicMock() + wallet.request_mint = AsyncMock(side_effect=httpx.ReadTimeout("response lost")) + + with ( + patch.object(settings, "primary_mint", TRUSTED), + patch.object(settings, "cashu_mints", [TRUSTED]), + patch.object(settings, "mint_retry_max_attempts", 3), + patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)), + patch("routstr.lightning.mint_cooldown_remaining", return_value=0.0), + pytest.raises(Exception), + ): + await _request_mint_with_fallback(10) + + wallet.request_mint.assert_awaited_once_with(10) + + async def test_trusted_allowed_mints_are_used_verbatim() -> None: attempted: list[str] = [] diff --git a/tests/unit/test_ppq_resilience.py b/tests/unit/test_ppq_resilience.py new file mode 100644 index 00000000..b1dd9f9c --- /dev/null +++ b/tests/unit/test_ppq_resilience.py @@ -0,0 +1,103 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from routstr.upstream.ppqai import ( + PPQAIUpstreamProvider, + PPQCircuitOpenError, + _ppq_circuits, + _safe_read_request, +) + + +@pytest.fixture(autouse=True) +def _clear_ppq_circuits() -> None: + _ppq_circuits.clear() + + +@pytest.mark.asyncio +async def test_safe_ppq_read_retries_timeout_with_bounded_backoff() -> None: + request = httpx.Request("GET", "https://api.ppq.ai/models") + success = httpx.Response(200, request=request, json={"data": []}) + client = MagicMock() + client.request = AsyncMock( + side_effect=[httpx.ReadTimeout("", request=request), success] + ) + with ( + patch("routstr.upstream.ppqai.random.uniform", return_value=0.0), + patch("routstr.upstream.ppqai.asyncio.sleep", AsyncMock()) as sleep, + ): + response = await _safe_read_request( + client, "GET", "https://api.ppq.ai/models", headers={} + ) + + assert response is success + assert client.request.await_count == 2 + sleep.assert_awaited_once_with(0.25) + + +@pytest.mark.asyncio +async def test_safe_ppq_read_opens_cross_cycle_circuit_and_probe_clears_it() -> None: + request = httpx.Request("GET", "https://api.ppq.ai/models") + client = MagicMock() + client.request = AsyncMock(side_effect=httpx.ReadTimeout("down", request=request)) + + with ( + patch("routstr.upstream.ppqai.random.uniform", return_value=0.0), + patch("routstr.upstream.ppqai.asyncio.sleep", AsyncMock()), + patch("routstr.upstream.ppqai.time.monotonic", return_value=100.0), + pytest.raises(httpx.ReadTimeout), + ): + await _safe_read_request(client, "GET", str(request.url), headers={}) + assert client.request.await_count == 3 + + with ( + patch("routstr.upstream.ppqai.time.monotonic", return_value=110.0), + pytest.raises(PPQCircuitOpenError), + ): + await _safe_read_request(client, "GET", str(request.url), headers={}) + assert client.request.await_count == 3 + + success = httpx.Response(200, request=request, json={"data": []}) + client.request = AsyncMock(return_value=success) + with patch("routstr.upstream.ppqai.time.monotonic", return_value=131.0): + assert ( + await _safe_read_request(client, "GET", str(request.url), headers={}) + is success + ) + + state = next(iter(_ppq_circuits.values())) + assert state.consecutive_failures == 0 + assert state.cooldown_until == 0.0 + + +@pytest.mark.asyncio +async def test_fetch_models_failure_is_not_a_valid_empty_catalog() -> None: + provider = PPQAIUpstreamProvider("secret") + with ( + patch( + "routstr.upstream.ppqai._safe_read_request", + AsyncMock(side_effect=httpx.ReadTimeout("catalog timed out")), + ), + pytest.raises(httpx.ReadTimeout), + ): + await provider.fetch_models() + + +@pytest.mark.asyncio +async def test_ppq_invoice_creation_post_is_never_retried() -> None: + provider = PPQAIUpstreamProvider("secret") + client = MagicMock() + client.post = AsyncMock(side_effect=httpx.ReadTimeout("invoice timed out")) + context = MagicMock() + context.__aenter__ = AsyncMock(return_value=client) + context.__aexit__ = AsyncMock(return_value=None) + + with ( + patch("routstr.upstream.ppqai.httpx.AsyncClient", return_value=context), + pytest.raises(httpx.ReadTimeout), + ): + await provider.create_lightning_topup(10, "USD") + + client.post.assert_awaited_once() diff --git a/tests/unit/test_refund_no_retry.py b/tests/unit/test_refund_no_retry.py new file mode 100644 index 00000000..26ab0916 --- /dev/null +++ b/tests/unit/test_refund_no_retry.py @@ -0,0 +1,47 @@ +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from fastapi import HTTPException + +from routstr.upstream.base import BaseUpstreamProvider + + +@pytest.mark.asyncio +async def test_send_refund_does_not_retry_ambiguous_token_creation() -> None: + provider = object.__new__(BaseUpstreamProvider) + send_token = AsyncMock(side_effect=httpx.ReadTimeout("swap response lost")) + store = AsyncMock() + + with ( + patch("routstr.upstream.base.send_token", send_token), + patch("routstr.upstream.base.store_cashu_transaction", store), + pytest.raises(HTTPException) as raised, + ): + await provider.send_refund(10, "sat", mint="https://mint.test") + + assert raised.value.status_code == 401 + send_token.assert_awaited_once_with( + 10, unit="sat", mint_url="https://mint.test" + ) + store.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_send_refund_does_not_retry_or_store_on_generic_failure() -> None: + provider = object.__new__(BaseUpstreamProvider) + send_token = AsyncMock(side_effect=Exception("mint rejected swap")) + store = AsyncMock() + + with ( + patch("routstr.upstream.base.send_token", send_token), + patch("routstr.upstream.base.store_cashu_transaction", store), + pytest.raises(HTTPException) as raised, + ): + await provider.send_refund(10, "sat", mint="https://mint.test") + + assert raised.value.status_code == 401 + send_token.assert_awaited_once_with( + 10, unit="sat", mint_url="https://mint.test" + ) + store.assert_not_awaited() diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 94ef65a4..ea8b6b1f 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -4,6 +4,7 @@ from collections.abc import AsyncGenerator from typing import cast from unittest.mock import AsyncMock, MagicMock, patch +import httpx import pytest from fastapi import BackgroundTasks from sqlalchemy.exc import SQLAlchemyError @@ -340,6 +341,153 @@ async def test_responses_streaming_releases_and_raises_on_billing_failure( release.assert_awaited_once_with(snapshot, session, 500) +@pytest.mark.asyncio +@pytest.mark.parametrize("api", ["chat", "responses"]) +@pytest.mark.parametrize("finalization_fails", [False, True]) +async def test_partial_remote_protocol_error_finalizes_and_closes_once( + api: str, + finalization_fails: bool, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n' + raise httpx.RemoteProtocolError("incomplete chunked read") + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aiter_bytes = aiter_bytes + upstream_response.aclose = AsyncMock() + client = MagicMock() + client.aclose = AsyncMock() + key = MagicMock(spec=ApiKey) + key.hashed_key = f"{api}-partial" + key.balance = 10_000 + session = MagicMock() + session.get = AsyncMock(return_value=key) + session.rollback = AsyncMock() + session_context = MagicMock() + session_context.__aenter__ = AsyncMock(return_value=session) + session_context.__aexit__ = AsyncMock(return_value=None) + adjust = ( + AsyncMock(side_effect=SQLAlchemyError("database unavailable")) + if finalization_fails + else AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0}) + ) + snapshot = ReservationSnapshot( + release_id=f"{api}-partial-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + release = AsyncMock(return_value=True) + + with ( + patch("routstr.upstream.base.adjust_payment_for_tokens", adjust), + patch("routstr.upstream.base.release_reservation", release), + patch("routstr.upstream.base.create_session", return_value=session_context), + ): + if api == "chat": + response = await provider.handle_streaming_chat_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + background_tasks=BackgroundTasks(), + reservation_snapshot=snapshot, + client=client, + ) + else: + response = await provider.handle_streaming_responses_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + client=client, + ) + emitted = bytearray() + with pytest.raises(httpx.RemoteProtocolError): + async for chunk in response.body_iterator: + emitted.extend( + chunk.encode() if isinstance(chunk, str) else bytes(chunk) + ) + + adjust.assert_awaited_once() + if finalization_fails: + session.rollback.assert_awaited_once() + release.assert_awaited_once_with(snapshot, session, 500) + else: + release.assert_not_awaited() + upstream_response.aclose.assert_awaited_once() + client.aclose.assert_awaited_once() + assert b"[DONE]" not in emitted + + +@pytest.mark.asyncio +@pytest.mark.parametrize("api", ["chat", "responses"]) +async def test_partial_stream_preserves_transport_error_when_billing_db_is_down( + api: str, +) -> None: + provider = BaseUpstreamProvider( + base_url="https://api.example.com", api_key="test-key" + ) + + async def aiter_bytes() -> AsyncGenerator[bytes, None]: + yield b'data: {"model":"test","choices":[]}\n\n' + raise httpx.RemoteProtocolError("incomplete chunked read") + + upstream_response = MagicMock( + status_code=200, headers={"content-type": "text/event-stream"} + ) + upstream_response.aiter_bytes = aiter_bytes + upstream_response.aclose = AsyncMock() + client = MagicMock() + client.aclose = AsyncMock() + key = MagicMock(spec=ApiKey) + key.hashed_key = f"{api}-database-down" + key.balance = 10_000 + snapshot = ReservationSnapshot( + release_id=f"{api}-database-down-release", + key_hash=key.hashed_key, + billing_key_hash=key.hashed_key, + reserved_msats=500, + ) + unavailable_session = MagicMock() + unavailable_session.__aenter__ = AsyncMock( + side_effect=SQLAlchemyError("database unavailable") + ) + unavailable_session.__aexit__ = AsyncMock(return_value=None) + + with patch( + "routstr.upstream.base.create_session", return_value=unavailable_session + ): + if api == "chat": + response = await provider.handle_streaming_chat_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + background_tasks=BackgroundTasks(), + reservation_snapshot=snapshot, + client=client, + ) + else: + response = await provider.handle_streaming_responses_completion( + response=upstream_response, + key=key, + max_cost_for_model=500, + reservation_snapshot=snapshot, + client=client, + ) + with pytest.raises(httpx.RemoteProtocolError, match="incomplete chunked read"): + async for _ in response.body_iterator: + pass + + upstream_response.aclose.assert_awaited_once() + client.aclose.assert_awaited_once() + + @pytest.mark.asyncio async def test_responses_streaming_duplicate_publishes_zero_settled_cost() -> None: provider = BaseUpstreamProvider( @@ -578,9 +726,7 @@ async def test_client_disconnect_midstream_finalizes_and_stops_heartbeat() -> No background_tasks=background_tasks, reservation_snapshot=snapshot, ) - iterator = cast( - AsyncGenerator[bytes, None], response.body_iterator - ) + iterator = cast(AsyncGenerator[bytes, None], response.body_iterator) await iterator.__anext__() # first chunk reaches the client await iterator.aclose() # client aborts the socket here diff --git a/tests/unit/test_wallet.py b/tests/unit/test_wallet.py index 7c5de1fe..5ad0a4f3 100644 --- a/tests/unit/test_wallet.py +++ b/tests/unit/test_wallet.py @@ -26,6 +26,7 @@ from routstr.wallet import ( recieve_token, send, send_token, + send_token_from_owner_locked, ) @@ -40,13 +41,19 @@ def isolate_wallet_runtime_state() -> Generator[None, None, None]: wallet_module._MintRateGuard._guards.clear() wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() + wallet_module._wallet_last_mint_load.clear() wallet_module._wallet_load_locks.clear() + wallet_module._mint_metadata_last_load.clear() + wallet_module._mint_metadata_load_locks.clear() yield settings.mint_max_concurrency = original_concurrency wallet_module._MintRateGuard._guards.clear() wallet_module._wallets.clear() wallet_module._wallet_last_load.clear() + wallet_module._wallet_last_mint_load.clear() wallet_module._wallet_load_locks.clear() + wallet_module._mint_metadata_last_load.clear() + wallet_module._mint_metadata_load_locks.clear() @pytest.mark.asyncio @@ -66,6 +73,63 @@ async def test_get_balance() -> None: assert balance == 50000 +@pytest.mark.asyncio +async def test_wallet_metadata_is_reused_across_units() -> None: + from routstr.wallet import Wallet + + sat_wallet = MagicMock(url="http://mint:3338") + sat_wallet.load_mint_keysets = AsyncMock() + sat_wallet.activate_keyset = AsyncMock() + sat_wallet.load_mint_info = AsyncMock() + sat_wallet.load_keysets_from_db = AsyncMock() + + msat_wallet = MagicMock(url="http://mint:3338") + msat_wallet.load_mint_keysets = AsyncMock() + msat_wallet.activate_keyset = AsyncMock() + msat_wallet.load_mint_info = AsyncMock() + msat_wallet.load_keysets_from_db = AsyncMock() + + with patch("routstr.wallet.time.monotonic", return_value=1000.0): + await Wallet.load_mint(sat_wallet) + await Wallet.load_mint(msat_wallet) + + sat_wallet.load_mint_keysets.assert_awaited_once_with(False) + sat_wallet.load_mint_info.assert_awaited_once_with(reload=True) + msat_wallet.load_mint_keysets.assert_not_awaited() + msat_wallet.load_keysets_from_db.assert_awaited_once_with() + msat_wallet.load_mint_info.assert_awaited_once_with(reload=False) + + +@pytest.mark.asyncio +async def test_get_wallet_refreshes_local_proofs_without_reloading_mint() -> None: + from routstr import wallet as wallet_module + 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)), + patch("routstr.wallet.time.monotonic", return_value=1000.0), + ): + await get_wallet("http://mint:3338", "sat") + wallet_module._wallet_last_load["http://mint:3338_sat"] = 900.0 + await get_wallet("http://mint:3338", "sat") + + assert mock_wallet.load_mint.await_count == 1 + assert mock_wallet.load_proofs.await_count == 2 + + +@pytest.mark.asyncio +async def test_get_wallet_quote_only_skips_proof_reload() -> 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", load_proofs=False) + + mock_wallet.load_mint.assert_awaited_once_with() + mock_wallet.load_proofs.assert_not_awaited() + + @pytest.mark.asyncio async def test_get_wallet_force_reload_bypasses_reload_interval() -> None: from routstr.wallet import get_wallet @@ -392,6 +456,33 @@ async def test_send_token() -> None: assert token == "test_token" +@pytest.mark.asyncio +async def test_owner_only_token_rejects_customer_backed_proofs() -> None: + mint = "http://mint:3338" + proof = Mock(amount=1000, reserved=False) + wallet = Mock(keysets={}, proofs=[proof], select_to_send=AsyncMock()) + + with ( + patch( + "routstr.wallet.find_trusted_mint_with_funds", + AsyncMock(return_value=mint), + ), + patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)), + patch( + "routstr.wallet.get_proofs_per_mint_and_unit", + return_value=[proof], + ), + patch( + "routstr.wallet._owner_balance_for_mint_and_unit", + AsyncMock(return_value=50), + ), + pytest.raises(ValueError, match="Owner Cashu balance"), + ): + await send_token_from_owner_locked(100, "sat", mint) + + wallet.select_to_send.assert_not_awaited() + + @pytest.mark.asyncio async def test_release_token_reservation_unreserves_local_proofs() -> None: from routstr.wallet import release_token_reservation @@ -632,6 +723,38 @@ async def test_credit_balance() -> None: assert mock_session.refresh.called +@pytest.mark.asyncio +async def test_concurrent_duplicate_token_credits_exactly_once() -> None: + key = Mock(balance=0, hashed_key="duplicate-key") + session = AsyncMock() + session.exec.return_value.rowcount = 1 + session.refresh = AsyncMock() + receive = AsyncMock( + side_effect=[ + (1000, "sat", "https://mint.test"), + ValueError("Mint Error: proofs already spent (Code: 11001)"), + ] + ) + store = AsyncMock() + + with ( + patch("routstr.wallet.recieve_token", receive), + patch("routstr.wallet.store_cashu_transaction", store), + ): + results = await asyncio.gather( + credit_balance("cashuAduplicate", key, session), + credit_balance("cashuAduplicate", key, session), + return_exceptions=True, + ) + + assert sum(result == 1_000_000 for result in results) == 1 + failure = next(result for result in results if isinstance(result, Exception)) + classified = classify_redemption_error(failure) + assert classified is not None and classified[3] == "cashu_token_already_spent" + assert session.exec.await_count == 1 + store.assert_awaited_once() + + @pytest.mark.asyncio async def test_credit_balance_constrains_redemption_to_key_mint() -> None: key_mint = "http://key-mint:3338" @@ -2970,6 +3093,7 @@ async def test_load_mint_propagates_rate_limit() -> None: from routstr.wallet import Wallet wallet = Wallet.__new__(Wallet) + wallet.url = "https://rate-limited-mint.example" error = MintRateLimitedError( "Cashu mint rate limited", request=httpx.Request("GET", "https://mint.example/v1/keysets"), @@ -2987,6 +3111,7 @@ async def test_load_mint_propagates_connection_error() -> None: from routstr.wallet import Wallet wallet = Wallet.__new__(Wallet) + wallet.url = "https://unavailable-mint.example" error = httpx.ConnectError("mint unavailable") with ( patch.object(wallet, "load_mint_keysets", new=AsyncMock(side_effect=error)), @@ -3002,6 +3127,7 @@ async def test_load_mint_runs_keysets_activation_and_info() -> None: from routstr.wallet import Wallet wallet = Wallet.__new__(Wallet) + wallet.url = "https://mint-load.example" with ( patch.object(wallet, "load_mint_keysets", new=AsyncMock()) as load_keysets, patch.object(wallet, "activate_keyset", new=AsyncMock()) as activate,