fix: harden melt sizing, mint cooldowns, PPQ reads, and streaming finalization

This commit is contained in:
9qeklajc
2026-08-26 09:59:16 +02:00
parent aae2c9059e
commit ba9cecc916
24 changed files with 1391 additions and 258 deletions
+2
View File
@@ -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.
+21
View File
@@ -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.
+26 -5
View File
@@ -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,
}
},
)
+14 -2
View File
@@ -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,
},
)
+24 -1
View File
@@ -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)
+98 -31
View File
@@ -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(
+162 -81
View File
@@ -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.
+109 -52
View File
@@ -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]:
+101 -12
View File
@@ -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(
+38 -4
View File
@@ -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
+10 -1
View File
@@ -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")
+36 -2
View File
@@ -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
@@ -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(
+18 -20
View File
@@ -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
+102 -10
View File
@@ -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)
+86 -1
View File
@@ -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()
+8 -4
View File
@@ -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},
+9 -1
View File
@@ -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",
+68 -12
View File
@@ -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
+87 -14
View File
@@ -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)
+29
View File
@@ -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
+103
View File
@@ -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()
@@ -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
+60
View File
@@ -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"