From ba9cecc91683741c9c8ea7f70a2b6c7ea4c2c661 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 26 Aug 2026 09:59:16 +0200 Subject: [PATCH] fix: harden melt sizing, mint cooldowns, PPQ reads, and streaming finalization --- .env.example | 2 + docs/adr/001-cashu-payment-safety.md | 21 ++ routstr/auth.py | 31 ++- routstr/core/exceptions.py | 16 +- routstr/mint.py | 25 +- routstr/payment/lnurl.py | 129 +++++++--- routstr/upstream/auto_topup.py | 243 ++++++++++++------ routstr/upstream/base.py | 161 ++++++++---- routstr/upstream/ppqai.py | 113 +++++++- routstr/wallet.py | 42 ++- tests/integration/conftest.py | 11 +- .../integration/test_ppq_auto_topup_claim.py | 38 ++- .../test_routstr_auto_topup_claim.py | 33 ++- .../integration/test_wallet_authentication.py | 38 ++- tests/integration/test_wallet_melt_restart.py | 112 +++++++- tests/unit/test_auto_topup.py | 87 ++++++- tests/unit/test_core_exceptions.py | 12 +- tests/unit/test_lightning_settlement.py | 10 +- .../unit/test_lnurl_amount_and_destination.py | 80 +++++- tests/unit/test_lnurl_melt_timeout.py | 101 +++++++- tests/unit/test_mint.py | 29 +++ tests/unit/test_ppq_resilience.py | 103 ++++++++ .../test_streaming_billing_finalization.py | 152 ++++++++++- tests/unit/test_wallet.py | 60 +++++ 24 files changed, 1391 insertions(+), 258 deletions(-) create mode 100644 docs/adr/001-cashu-payment-safety.md create mode 100644 tests/unit/test_ppq_resilience.py diff --git a/.env.example b/.env.example index 35a171ba..5ad7a74d 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 +# Seconds a file-backed SQLite writer waits for the write lock (default: 30). +# 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/docs/adr/001-cashu-payment-safety.md b/docs/adr/001-cashu-payment-safety.md new file mode 100644 index 00000000..a2ed627f --- /dev/null +++ b/docs/adr/001-cashu-payment-safety.md @@ -0,0 +1,21 @@ +# ADR-001: Cashu payment safety boundaries + +## Status + +Accepted + +## Context + +Cashu proofs are bearer instruments. Retrying quote creation, refund delivery, or a dispatched Lightning melt can duplicate side effects or spend proofs whose outcome is still unknown. Mint transport failures and concurrent workers also need one shared policy. + +## Decision + +- Treat account, invoice, quote, melt, token-delivery, and refund creation as non-idempotent unless an upstream idempotency key is available. +- A dispatched melt with an unknown outcome keeps a durable quote-linked proof reservation until later reconciliation confirms a terminal state. An immediate `unpaid` observation after transport loss is not terminal. +- Size melts from the quote amount, reserve, and exact proof input fees within the caller's gross budget; do not use recursive send selection for melt planning. +- Apply mint transport/rate cooldowns centrally and permit only explicit reconciliation probes during cooldown. +- Auto-topups require fresh threshold confirmation, durable per-provider claims/cooldown, atomic spend-cap checks, and owner-only funds. + +## Consequences + +Transient failures can delay payouts/topups rather than risk duplicate payment. Operators may need to reconcile ambiguous claims. Tests must cover restart, concurrency, and partial-stream failures at these boundaries. 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/exceptions.py b/routstr/core/exceptions.py index 9e4f2ae6..7e319a94 100644 --- a/routstr/core/exceptions.py +++ b/routstr/core/exceptions.py @@ -40,15 +40,27 @@ 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. + # Retryable mint outages are expected dependency failures, not application + # faults, so keep them visible without flooding the error stream. 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/mint.py b/routstr/mint.py index 64632b9e..8948bd4f 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,18 @@ def mint_cooldown_reason(mint_url: str) -> str | None: return MintRateGuard.get(mint_url).cooldown_reason() +def is_mint_transport_error(error: BaseException) -> bool: + """Return whether an exception chain contains a mint transport failure.""" + 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 +283,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 +297,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() @@ -292,6 +307,10 @@ async def run_mint_operation( except MintCooldownError: raise except (asyncio.TimeoutError, httpx.TimeoutException) as exc: + if guard is not None: + guard.apply_cooldown( + MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport" + ) if retry_timeouts and attempt < max_attempts - 1: backoff = (2**attempt) + (time.monotonic() % 1.0) logger.warning( @@ -310,6 +329,10 @@ async def run_mint_operation( ) 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 814f83fa..3e391f8b 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 @@ -238,6 +237,35 @@ 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 required <= gross_budget and selected_amount >= required: + return selected, 0 + return None, max(1, required - min(selected_amount, gross_budget)) + + async def raw_send_to_lnurl( wallet: Wallet, proofs: list[Proof], @@ -293,38 +321,55 @@ 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 + # Start at the caller's gross budget and converge downward from the mint's + # exact reserve plus NUT-02 input fees. Starting below the budget with a + # percentage heuristic silently underpays even when the exact fees are tiny. + 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), - # Creating another quote after response loss only abandons the first. - retry_timeouts=False, - ) - - # 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 + # Fee reserves can change with the invoice amount. Each quote reduces the + # candidate by its exact shortfall, so this bounded fixed-point search keeps + # the largest amount the gross budget can fund without Cashu coin selection. + 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), + # Creating another quote after response loss only abandons the first. + 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) + proofs = selected_proofs + await wallet.set_reserved_for_send(proofs, reserved=True) try: melt_response = await run_mint_operation( @@ -365,8 +410,14 @@ async def raw_send_to_lnurl( 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: + # A direct unpaid response is authoritative: the mint rejected the melt + # and Cashu has already cleared its quote-linked reservation. + 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( @@ -374,6 +425,9 @@ async def raw_send_to_lnurl( op_name="reconcile_lnurl_melt_quote", mint_url=str(wallet.url), retry_timeouts=False, + # One direct state lookup is required to reconcile the just-dispatched + # melt even though its transport failure opened the mint cooldown. + allow_during_cooldown=True, ) except Exception as reconciliation_error: raise MeltOutcomeAmbiguousError( @@ -384,9 +438,22 @@ 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: - # get_melt_quote() has authoritatively released the melt reservation; - # callers may restore their debit and retry with a new payment plan. - raise LNURLError("Cashu mint confirmed that the melt was unpaid") from melt_error + # Reaching reconciliation means melt was dispatched and either lost its + # response or returned pending. A just-dispatched quote can briefly read + # UNPAID before transitioning. Cashu clears the reservation while + # refreshing that state, so restore it and require later reconciliation. + 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..64d7d380 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,64 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None: ) try: - token = await send_token(amount, "sat", mint_url) + async with wallet_operation_guard(): + # The cap, owner-liability check, proof reservation, and outgoing + # audit row share one wallet mutation scope. The audit row must be + # durable before another worker can recheck the rolling cap. + 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, - ) - + # The audit row and SENT claim committed together before this network call, + # so a worker crash cannot make reconciliation treat reserved proofs as an + # unspent CLAIMED attempt. result = await provider.topup(token) if "error" in result: @@ -705,7 +704,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 +713,60 @@ 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") + """Commit the bearer-token audit row and SENT claim atomically.""" + 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 +1120,13 @@ 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, + # For successful payments this timestamps the durable cooldown, + # not merely when the original claim was created. + created_at=int(time.time()) if collected else CashuTransaction.created_at, + ) ) updated = (getattr(result, "rowcount", 0) or 0) == 1 if updated: @@ -1106,8 +1151,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 +1220,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 +1231,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 +1251,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 +1292,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 +1513,27 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None if balance >= threshold_usd: return + # A single stale/partial balance response must never create an invoice. + # Read the uncached endpoint again and require independent agreement. + 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 595ec670..e0f0a4f9 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import inspect import json import math import traceback @@ -70,6 +71,17 @@ 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 + + CostMetadata = CostData | MaxCostData | dict[str, Any] @@ -993,6 +1005,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 +1047,42 @@ 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, + try: + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key ) - usage_finalized = True - except Exception: - pass + 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: + # Preserve the original stream exception. If the database + # cannot even be opened/read, stale-reservation cleanup is + # the only safe recovery path. + 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 +1312,23 @@ 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) + try: + if not usage_finalized: + await finalize_db_only() + finally: + try: + await _aclose_if_needed(response) + finally: + await _aclose_if_needed(client) # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -1451,6 +1488,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 +1522,42 @@ 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, + try: + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key ) - usage_finalized = True - except Exception: - pass + 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: + # Preserve the original stream exception. If the database + # cannot even be opened/read, stale-reservation cleanup is + # the only safe recovery path. + 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 +1747,23 @@ 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() + try: + if not usage_finalized: + await finalize_db_only() + finally: + try: + await _aclose_if_needed(response) + finally: + await _aclose_if_needed(client) # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) @@ -3051,9 +3115,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 +3123,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 +3393,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: @@ -5339,7 +5396,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..d373d8e8 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,90 @@ if TYPE_CHECKING: logger = get_logger(__name__) +_PPQ_SAFE_READ_ATTEMPTS = 3 +_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0 + + +class PPQCircuitOpenError(RuntimeError): + """PPQ safe reads are suppressed until one probe is allowed.""" + + +@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) + return f"{parsed.scheme}://{parsed.host}:{parsed.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: + """Retry safe reads, then open one process-local circuit per PPQ origin.""" + state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState()) + loop = asyncio.get_running_loop() + if state.loop is not loop: + # Runtime uses one long-lived loop; pytest and some embedded hosts do + # not. Preserve circuit state while replacing a loop-bound lock. + 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 @@ -125,8 +213,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): try: 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 + ) data = response.json() models_data = data.get("data", []) @@ -233,12 +322,10 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider): return models - except Exception as e: - logger.error( - "Error fetching models from PPQ.AI", - extra={"error": str(e), "error_type": type(e).__name__}, - ) - return [] + except Exception: + # The base refresh handler preserves the last good model cache when + # fetching raises; [] would look like a valid empty catalog. + raise async def on_upstream_error_redirect( self, status_code: int, error_message: str @@ -360,8 +447,9 @@ 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 +548,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 3814ac42..24072c74 100644 --- a/routstr/wallet.py +++ b/routstr/wallet.py @@ -524,7 +524,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 @@ -534,6 +538,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) @@ -580,6 +590,14 @@ 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: + """Create an owner-funded token while the caller holds the wallet guard.""" + _, 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. @@ -1824,9 +1842,25 @@ 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 diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index aa10a81c..8dcd1876 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,15 @@ 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]: + """Do not let one integration test's simulated outage poison the next.""" + 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..82c33c6f 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,15 @@ 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. + # The durable cooldown blocks an immediate duplicate, then the claim lock + # can be reused without overwriting audit history after it expires. + 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_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..2054f0af 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,35 @@ 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""" + """Concurrent duplicate redemption credits once and deterministically replays.""" - # 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 + # The real DB contains one logical credit. A later replay is also mutation-free. + 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..6c28963c 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,35 @@ 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: + """Actual Wallet.get_melt_quote consumes quote-linked proofs after restart.""" + 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 +143,65 @@ 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.""" + """Actual Wallet.get_melt_quote releases proofs after an unpaid answer.""" 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_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..53cb13df 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,29 @@ LNURL_DATA = { "max_sendable": 100_000_000, } -# 1000 sat minus the 11 sat estimated fee reserve. -EXPECTED_QUOTE_SAT = 989 +# The exact plan spends the 1000 sat gross budget as 999 sat + 1 sat reserve. +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 +135,22 @@ 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) + # 1_000_000 msat gross minus the exact 1 msat reserve leaves 999_999 msat. + 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 +158,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 diff --git a/tests/unit/test_lnurl_melt_timeout.py b/tests/unit/test_lnurl_melt_timeout.py index dd0dfb02..c517cbe2 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 @@ -16,6 +17,14 @@ from routstr.payment.lnurl import ( 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, @@ -23,18 +32,22 @@ LNURL_DATA = { } -# 1000 sat minus the 11 sat estimated fee reserve. -QUOTE_AMOUNT_SAT = 989 +# Exact planning pays 999 sat from a 1000 sat budget with a 1 sat reserve. +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 @@ -54,7 +67,26 @@ def _lnurl_patches() -> tuple[Any, Any]: @pytest.mark.asyncio -async def test_raw_send_to_lnurl_timeout_reconciled_unpaid_is_retry_safe() -> 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: @@ -71,19 +103,19 @@ async def test_raw_send_to_lnurl_timeout_reconciled_unpaid_is_retry_safe() -> No patch.object(settings, "mint_retry_max_attempts", 0), data_patch, invoice_patch, - pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised, + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), ): await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) - assert not isinstance(raised.value, MeltOutcomeAmbiguousError) wallet.get_melt_quote.assert_awaited_once_with("q") - wallet.set_reserved_for_melt.assert_awaited_once_with( + 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_error_is_reconciled_once() -> None: +async def test_raw_send_to_lnurl_wrapped_transport_unpaid_remains_ambiguous() -> None: wallet, proofs = _wallet() async def _wrapped_transport_error(**kwargs: object) -> None: @@ -102,14 +134,14 @@ async def test_raw_send_to_lnurl_wrapped_transport_error_is_reconciled_once() -> patch.object(settings, "mint_retry_max_attempts", 3), data_patch, invoice_patch, - pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised, + pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"), ): await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000) - assert not isinstance(raised.value, MeltOutcomeAmbiguousError) wallet.melt.assert_awaited_once() wallet.get_melt_quote.assert_awaited_once_with("q") - wallet.set_reserved_for_melt.assert_awaited_once_with( + assert wallet.set_reserved_for_melt.await_count == 2 + wallet.set_reserved_for_melt.assert_awaited_with( proofs, reserved=True, quote_id="q" ) @@ -180,6 +212,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( @@ -213,7 +284,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 @@ -238,7 +310,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..dc58d10c 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,34 @@ 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) + + async def test_guard_concurrency_change_preserves_cooldown_state() -> None: from routstr.core.settings import settings 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_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 8b738f89..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, ) @@ -455,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 @@ -695,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"