Merge pull request #693 from Routstr/fix/cashu-mint-interactions

fix: harden wallet and Cashu operations
This commit is contained in:
9qeklajc
2026-08-26 23:28:08 +02:00
committed by GitHub
32 changed files with 1872 additions and 482 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
# SQLite write-lock timeout, in seconds.
# DATABASE_BUSY_TIMEOUT=30
# SQLite serialises writes; increasing its pool can trade pool timeouts for
# "database is locked" errors rather than increasing write throughput.
+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,
}
},
)
+4 -18
View File
@@ -1,4 +1,3 @@
import asyncio
import json
import re
import secrets
@@ -1415,12 +1414,7 @@ async def initiate_provider_topup(
else {}
)
last_status_code = 500
last_error_detail: object = "Failed to create top-up invoice"
# Some upstream Routstr nodes fail the first invoice request after warm-up
# and succeed immediately on retry. Retry once here so the UI stays single-click.
for attempt in range(2):
# Quote creation is unsafe to retry without idempotency.
resp = await client.post(
f"{clean_url}/v1/balance/lightning/invoice",
json=request_json,
@@ -1442,23 +1436,15 @@ async def initiate_provider_topup(
f"Upstream topup request failed: {resp.text}",
extra={
"provider_id": provider_id,
"attempt": attempt + 1,
"status_code": resp.status_code,
},
)
try:
last_error_detail = resp.json()
error_detail: object = resp.json()
except Exception:
last_error_detail = resp.text
last_status_code = resp.status_code
if resp.status_code < 500 or attempt == 1:
break
await asyncio.sleep(0.2)
error_detail = resp.text
raise HTTPException(
status_code=last_status_code, detail=last_error_detail
status_code=resp.status_code, detail=error_detail
)
upstream_instance = _instantiate_provider(provider)
+7 -1
View File
@@ -37,6 +37,9 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"}
pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite
options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping}
connect_args: dict[str, object] = {}
if is_sqlite and not is_memory_sqlite:
connect_args["timeout"] = settings.database_busy_timeout
if not is_memory_sqlite:
options.update(
pool_size=settings.database_pool_size,
@@ -51,9 +54,12 @@ def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
"database_url_backend": backend,
"in_memory_sqlite": is_memory_sqlite,
**options,
"connect_args": connect_args,
},
)
created_engine = create_async_engine(database_url, echo=False, **options)
created_engine = create_async_engine(
database_url, echo=False, connect_args=connect_args, **options
)
hold_warn_seconds = settings.database_pool_hold_warn_seconds
def record_pool_checkout(
+12 -2
View File
@@ -44,15 +44,25 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
path = request.url.path
# 4xx is client behaviour; the uvicorn access log already records it.
# Only 5xx warrants a server-side warning/error log here.
if status_code >= 500:
logger.error(
error_type = None
if isinstance(detail, dict):
error = detail.get("error")
if isinstance(error, dict):
error_type = error.get("type")
log = (
logger.warning
if error_type in {"mint_unreachable", "mint_rate_limited"}
else logger.error
)
log(
f"HTTP {status_code} on {path}: {detail}",
extra={
"request_id": request_id,
"status_code": status_code,
"detail": detail,
"path": path,
"error_type": error_type,
},
)
+4
View File
@@ -148,6 +148,9 @@ class Settings(BaseSettings):
database_pool_hold_warn_seconds: float = Field(
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
)
database_busy_timeout: float = Field(
default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT"
)
# Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL")
@@ -209,6 +212,7 @@ ENV_ONLY_FIELDS = frozenset(
"database_pool_recycle",
"database_pool_pre_ping",
"database_pool_hold_warn_seconds",
"database_busy_timeout",
}
)
+9 -2
View File
@@ -215,11 +215,18 @@ async def _request_mint_with_fallback(
)
continue
try:
wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False)
wallet = await get_wallet(
mint_url,
"sat",
retry_on_rate_limit=False,
load_proofs=False,
)
quote = await run_mint_operation(
lambda: wallet.request_mint(amount_sats),
op_name="request_mint_invoice",
mint_url=mint_url,
# Quote creation is unsafe to retry without idempotency.
retry_timeouts=False,
retry_on_rate_limit=False,
)
return quote.request, quote.quote, mint_url
@@ -471,7 +478,7 @@ async def check_invoice_payment(
await session.commit()
mint_url = settlement.mint_url or settings.primary_mint
wallet = await get_wallet(mint_url, "sat")
wallet = await get_wallet(mint_url, "sat", load_proofs=False)
try:
mint_status = await run_mint_operation(
lambda: wallet.get_mint_quote(settlement.payment_hash),
+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,17 @@ def mint_cooldown_reason(mint_url: str) -> str | None:
return MintRateGuard.get(mint_url).cooldown_reason()
def is_mint_transport_error(error: BaseException) -> bool:
current: BaseException | None = error
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
if isinstance(current, MINT_TRANSPORT_EXCEPTIONS):
return True
current = current.__cause__ or current.__context__
return False
def is_mint_rate_limited(error: BaseException) -> bool:
"""Return whether an exception chain represents HTTP 429/cooldown."""
@@ -269,6 +282,7 @@ async def run_mint_operation(
mint_url: str = "",
retry_timeouts: bool = True,
retry_on_rate_limit: bool = True,
allow_during_cooldown: bool = False,
) -> Any:
"""Run one mint operation with bounded concurrency and adaptive cooldown."""
@@ -282,7 +296,7 @@ async def run_mint_operation(
return await factory()
async def invoke() -> Any:
if guard is not None:
if guard is not None and not allow_during_cooldown:
return await guard.run(timed_factory)
return await timed_factory()
@@ -293,6 +307,7 @@ async def run_mint_operation(
raise
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
if retry_timeouts and attempt < max_attempts - 1:
# Cooldown opens only after retries; earlier would stretch each backoff to a full cooldown wait.
backoff = (2**attempt) + (time.monotonic() % 1.0)
logger.warning(
"Mint operation timed out, retrying",
@@ -305,11 +320,19 @@ async def run_mint_operation(
)
await asyncio.sleep(backoff)
continue
if guard is not None:
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise httpx.TimeoutException(
f"{op_name} timed out (attempts: {attempt + 1})"
) from exc
except Exception as exc:
if not is_mint_rate_limited(exc):
if guard is not None and is_mint_transport_error(exc):
guard.apply_cooldown(
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
)
raise
backoff = (2**attempt) + (time.monotonic() % 1.0)
+89 -13
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
@@ -10,8 +9,8 @@ from cashu.core.base import MeltQuoteState
from cashu.wallet.wallet import Proof, Wallet
from ..mint import (
MINT_TRANSPORT_EXCEPTIONS,
is_mint_rate_limited,
is_mint_transport_error,
run_mint_operation,
)
@@ -226,6 +225,38 @@ async def get_lnurl_invoice(
return invoice_data["pr"], invoice_data
def _select_melt_proofs(
wallet: Wallet,
proofs: list[Proof],
*,
quote_amount: int,
fee_reserve: int,
gross_budget: int,
) -> tuple[list[Proof] | None, int]:
"""Select proofs that cover the quote and exact NUT-02 input fees.
Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a
wallet's full balance. Melts accept overpayment and return change, so a
bounded, largest-first selection is both safer and minimizes input fees.
"""
selected: list[Proof] = []
selected_amount = 0
required = quote_amount + fee_reserve
for proof in sorted(proofs, key=lambda item: item.amount, reverse=True):
if getattr(proof, "reserved", False) is True:
continue
selected.append(proof)
selected_amount += proof.amount
input_fees = int(wallet.get_fees_for_proofs(selected))
required = quote_amount + fee_reserve + input_fees
if selected_amount >= required:
if required <= gross_budget:
return selected, 0
# Covered but over budget; more proofs only raise input fees.
break
return None, max(1, required - min(selected_amount, gross_budget))
async def raw_send_to_lnurl(
wallet: Wallet,
proofs: list[Proof],
@@ -281,24 +312,24 @@ async def raw_send_to_lnurl(
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
)
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1
estimated_fees_msat = estimated_fees_sat * 1000
final_amount = amount_msat - estimated_fees_msat
final_amount = amount_msat
selected_proofs: list[Proof] | None = None
# Find the largest amount covered by the budget after reserve and input fees.
for _ in range(8):
if final_amount < lnurl_data["min_sendable"]:
raise LNURLError("Cashu melt fees leave no payable LNURL amount")
bolt11_invoice, _ = await get_lnurl_invoice(
lnurl_data["callback_url"], final_amount
)
melt_quote_resp = await run_mint_operation(
lambda: wallet.melt_quote(invoice=bolt11_invoice),
op_name="lnurl_melt_quote",
mint_url=str(wallet.url),
# Quote creation is unsafe to retry without idempotency.
retry_timeouts=False,
)
# 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:
@@ -307,10 +338,25 @@ async def raw_send_to_lnurl(
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
)
selected_proofs, shortfall = _select_melt_proofs(
wallet,
proofs,
quote_amount=quoted_amount,
fee_reserve=int(melt_quote_resp.fee_reserve),
gross_budget=amount,
)
if selected_proofs is not None:
break
final_amount -= shortfall * (1000 if unit == "sat" else 1)
else:
raise LNURLError("Cashu melt fees exceed the requested gross amount")
if on_melt_quote is not None:
await on_melt_quote(melt_quote_resp.quote)
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
assert selected_proofs is not None
proofs = selected_proofs
await wallet.set_reserved_for_send(proofs, reserved=True)
try:
melt_response = await run_mint_operation(
@@ -331,15 +377,29 @@ async def raw_send_to_lnurl(
# reserved as though a Lightning payment could still settle.
await wallet.set_reserved_for_send(proofs, reserved=False)
raise
if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS):
if not is_mint_transport_error(error):
raise
# Cashu clears reservations on transport errors despite an unknown outcome.
try:
await wallet.set_reserved_for_melt(
proofs, reserved=True, quote_id=melt_quote_resp.quote
)
except Exception as reservation_error:
raise MeltOutcomeAmbiguousError(
"Melt outcome is ambiguous and its proof reservation could not "
"be restored; proofs must not be retried"
) from reservation_error
melt_response = None
melt_error: BaseException | None = error
else:
melt_error = None
if getattr(melt_response, "state", None) == MeltQuoteState.paid:
melt_state = getattr(melt_response, "state", None)
if melt_state == MeltQuoteState.paid:
return final_amount
if melt_state == MeltQuoteState.unpaid:
await wallet.set_reserved_for_send(proofs, reserved=False)
raise LNURLError("Cashu mint confirmed that the melt was unpaid")
try:
quote = await run_mint_operation(
@@ -347,6 +407,8 @@ async def raw_send_to_lnurl(
op_name="reconcile_lnurl_melt_quote",
mint_url=str(wallet.url),
retry_timeouts=False,
# Reconciliation must bypass the cooldown opened by this failure.
allow_during_cooldown=True,
)
except Exception as reconciliation_error:
raise MeltOutcomeAmbiguousError(
@@ -356,6 +418,20 @@ async def raw_send_to_lnurl(
if quote is not None and quote.state == MeltQuoteState.paid:
return final_amount
if quote is not None and quote.state == MeltQuoteState.unpaid:
# A just-dispatched quote can briefly report unpaid before transitioning.
try:
await wallet.set_reserved_for_melt(
proofs, reserved=True, quote_id=melt_quote_resp.quote
)
except Exception as reservation_error:
raise MeltOutcomeAmbiguousError(
"Melt outcome is ambiguous and its proof reservation could not "
"be restored; proofs must not be retried"
) from reservation_error
raise MeltOutcomeAmbiguousError(
"Melt outcome is ambiguous; an immediate unpaid state is not final"
) from melt_error
state = getattr(getattr(quote, "state", None), "value", "unknown")
raise MeltOutcomeAmbiguousError(
+133 -61
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,34 +365,26 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
)
try:
token = await send_token(amount, "sat", mint_url)
except Exception as e:
logger.error(
"Failed to create cashu token for auto top-up",
extra={
"provider_id": row.id,
"amount": amount,
"mint_url": mint_url,
"error": str(e),
},
)
await _release_routstr_claim(row, operation_id)
return
async with wallet_operation_guard():
# Keep the spend cap and audit mutation in one wallet lock.
spent_24h_sats = await _routstr_spent_last_24h_sats()
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
raise ValueError("Routstr auto top-up daily spend cap reached")
token = await send_token_from_owner_locked(amount, "sat", mint_url)
actual_mint_url = token_mint_url(token, mint_url)
try:
await store_cashu_transaction(
await _persist_routstr_token_and_mark_sent(
row,
operation_id,
expected_sats=expected_sats,
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",
"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:
@@ -401,7 +395,7 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
extra={
"provider_id": row.id,
"mint_url": actual_mint_url,
"error": str(error),
"error": repr(error),
},
)
else:
@@ -409,21 +403,21 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
"Auto-topup token was released after persistence failed",
extra={"provider_id": row.id, "mint_url": actual_mint_url},
)
raise
except Exception as e:
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": repr(e),
"error_type": type(e).__name__,
},
)
await _release_routstr_claim(row, operation_id)
return
# Move the claim before the network call, not after: a worker that dies
# mid-request must leave behind a claim that says a token may already be
# with the peer.
await _mark_routstr_sent(
row,
operation_id,
expected_sats=expected_sats,
token=token,
amount=amount,
mint_url=actual_mint_url,
)
result = await provider.topup(token)
if "error" in result:
@@ -705,7 +699,7 @@ async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) ->
)
async def _mark_routstr_sent(
async def _persist_routstr_token_and_mark_sent(
row: UpstreamProviderRow,
operation_id: str,
*,
@@ -714,21 +708,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,
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:
async with create_session() as session:
@@ -1081,7 +1114,11 @@ async def _set_ppq_state_terminal(
col(CashuTransaction.collected) == False, # noqa: E712
col(CashuTransaction.swept) == False, # noqa: E712
)
.values(collected=collected, swept=swept)
.values(
collected=collected,
swept=swept,
created_at=int(time.time()) if collected else CashuTransaction.created_at,
)
)
updated = (getattr(result, "rowcount", 0) or 0) == 1
if updated:
@@ -1106,8 +1143,10 @@ async def _reconcile_ppq_state(
"""
async with create_session() as session:
transaction = await session.get(CashuTransaction, _ppq_state_id(row))
if transaction is None or transaction.collected or transaction.swept:
if transaction is None or transaction.swept:
return False
if transaction.collected:
return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS
claim = _parse_ppq_request_id(transaction.request_id)
if claim is None:
@@ -1173,7 +1212,7 @@ async def _reconcile_ppq_state(
async def _ppq_provider_is_claimable(
session: AsyncSession, provider_id: int | None
session: AsyncSession, row: UpstreamProviderRow
) -> bool:
"""Re-read the provider inside the claim transaction.
@@ -1184,10 +1223,16 @@ async def _ppq_provider_is_claimable(
this the worker could create a claim for a provider that no longer
exists, orphaning it forever.
"""
if provider_id is None:
if row.id is None:
return False
current = await session.get(UpstreamProviderRow, provider_id)
return current is not None and current.provider_type == "ppqai"
current = await session.get(UpstreamProviderRow, row.id)
return bool(
current is not None
and current.enabled
and current.provider_type == "ppqai"
and current.api_key == row.api_key
and current.provider_settings == row.provider_settings
)
async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
@@ -1198,10 +1243,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending")
async with create_session() as session:
if not await _ppq_provider_is_claimable(session, row.id):
if not await _ppq_provider_is_claimable(session, row):
return None
existing = await session.get(CashuTransaction, state_id)
if existing is not None:
if (
existing.collected
and not existing.swept
and int(time.time()) - existing.created_at
< PPQ_SETTLED_COOLDOWN_SECONDS
):
return None
result = await session.exec( # type: ignore[call-overload]
update(CashuTransaction)
.where(
@@ -1232,7 +1284,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
async with create_session() as session:
# Same fencing as the update path: the provider must still exist
# inside the transaction that creates the claim.
if not await _ppq_provider_is_claimable(session, row.id):
if not await _ppq_provider_is_claimable(session, row):
return None
session.add(
CashuTransaction(
@@ -1453,6 +1505,26 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None
if balance >= threshold_usd:
return
# Require two low-balance reads before creating an invoice.
confirmed_balance = await provider.get_balance()
if (
confirmed_balance is None
or not math.isfinite(confirmed_balance)
or confirmed_balance < 0
or confirmed_balance >= threshold_usd
):
logger.info(
"PPQ auto top-up aborted by balance confirmation",
extra={
"provider_id": row.id,
"first_balance_usd": balance,
"confirmed_balance_usd": confirmed_balance,
"threshold_usd": threshold_usd,
},
)
return
balance = confirmed_balance
# Perform local pricing and owner-funds checks before asking PPQ to create
# an invoice. The exact mint quote still has to be checked afterward, but
# predictable local failures should not leave abandoned PPQ invoices.
+94 -55
View File
@@ -1,12 +1,13 @@
from __future__ import annotations
import asyncio
import inspect
import json
import math
import traceback
import typing
import uuid
from collections.abc import AsyncGenerator, AsyncIterator, Iterator
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterator
from typing import Any, Mapping, Self, cast
import httpx
@@ -70,6 +71,32 @@ if typing.TYPE_CHECKING:
logger = get_logger(__name__)
async def _aclose_if_needed(resource: object | None) -> None:
if resource is None:
return
close = getattr(resource, "aclose", None)
if close is None:
return
result = close()
if inspect.isawaitable(result):
await result
async def _finalize_and_close_stream(
finalize: Callable[[], Awaitable[None]] | None,
response: object | None,
client: httpx.AsyncClient | None,
) -> None:
try:
if finalize is not None:
await finalize()
finally:
try:
await _aclose_if_needed(response)
finally:
await _aclose_if_needed(client)
CostMetadata = CostData | MaxCostData | dict[str, Any]
@@ -993,6 +1020,7 @@ class BaseUpstreamProvider:
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
client: httpx.AsyncClient | None = None,
) -> StreamingResponse:
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
@@ -1034,6 +1062,7 @@ class BaseUpstreamProvider:
nonlocal usage_finalized
if usage_finalized:
return
try:
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
@@ -1050,7 +1079,20 @@ class BaseUpstreamProvider:
)
usage_finalized = True
except Exception:
pass
logger.exception(
"Fallback stream billing finalization failed; releasing reservation",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key, new_session, reservation_snapshot
)
)
except Exception:
logger.exception(
"Fallback stream billing recovery could not access the database",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
def _process_event(
raw_event: bytes, final: bool = False
@@ -1280,18 +1322,24 @@ class BaseUpstreamProvider:
except Exception as stream_error:
logger.warning(
"Streaming interrupted; finalizing in background",
"Streaming interrupted; finalizing before closing upstream",
extra={
"error": str(stream_error),
"error_type": type(stream_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
finally:
if not usage_finalized:
# Create a background task to ensure finalization happens
# even if the generator is closed early
background_tasks.add_task(finalize_db_only)
# Shielded so a client disconnect cannot cancel billing
# finalization or leak the upstream connection.
await asyncio.shield(
_finalize_and_close_stream(
None if usage_finalized else finalize_db_only,
response,
client,
)
)
# Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers)
@@ -1451,6 +1499,7 @@ class BaseUpstreamProvider:
requested_model: str | None = None,
model_obj: Model | None = None,
reservation_snapshot: ReservationSnapshot | None = None,
client: httpx.AsyncClient | None = None,
) -> StreamingResponse:
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
@@ -1484,6 +1533,7 @@ class BaseUpstreamProvider:
nonlocal usage_finalized
if usage_finalized:
return
try:
async with create_session() as new_session:
fresh_key = await new_session.get(key.__class__, key.hashed_key)
if not fresh_key:
@@ -1500,7 +1550,20 @@ class BaseUpstreamProvider:
)
usage_finalized = True
except Exception:
pass
logger.exception(
"Fallback Responses billing finalization failed; releasing reservation",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
usage_finalized = (
await self._release_failed_streaming_reservation(
fresh_key, new_session, reservation_snapshot
)
)
except Exception:
logger.exception(
"Fallback Responses billing recovery could not access the database",
extra={"key_hash": key.hashed_key[:8] + "..."},
)
def _process_event(
raw_event: bytes, final: bool = False
@@ -1690,16 +1753,24 @@ class BaseUpstreamProvider:
except Exception as stream_error:
logger.warning(
"Responses API streaming interrupted; finalizing in background",
"Responses API streaming interrupted; finalizing before closing upstream",
extra={
"error": str(stream_error),
"error_type": type(stream_error).__name__,
"key_hash": key.hashed_key[:8] + "...",
},
)
raise
finally:
if not usage_finalized:
await finalize_db_only()
# Shielded so a client disconnect cannot cancel billing
# finalization or leak the upstream connection.
await asyncio.shield(
_finalize_and_close_stream(
None if usage_finalized else finalize_db_only,
response,
client,
)
)
# Remove inaccurate encoding headers from upstream response
response_headers = dict(response.headers)
@@ -3051,9 +3122,7 @@ class BaseUpstreamProvider:
if is_streaming and response.status_code == 200:
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
result = await self.handle_streaming_chat_completion(
return await self.handle_streaming_chat_completion(
response,
key,
max_cost_for_model,
@@ -3061,9 +3130,8 @@ class BaseUpstreamProvider:
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
client=client,
)
result.background = background_tasks
return result
# Handle both non-streaming chat completions and embeddings
if response.status_code == 200:
@@ -3332,19 +3400,15 @@ class BaseUpstreamProvider:
)
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_responses_completion(
return await self.handle_streaming_responses_completion(
response,
key,
max_cost_for_model,
requested_model=original_model_id,
model_obj=model_obj,
reservation_snapshot=reservation_snapshot,
client=client,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
background_tasks.add_task(client.aclose)
result.background = background_tasks
return result
if response.status_code == 200:
try:
@@ -3628,54 +3692,30 @@ class BaseUpstreamProvider:
extra={"amount": amount, "unit": unit, "mint": mint},
)
max_retries = 3
last_exception = None
refund_token = None
for attempt in range(max_retries):
try:
# Token creation may swap proofs, so it is unsafe to retry.
refund_token = await send_token(amount, unit=unit, mint_url=mint)
break
except Exception as e:
last_exception = e
if attempt < max_retries - 1:
logger.warning(
"Refund token creation failed, retrying",
extra={
"error": str(e),
"error_type": type(e).__name__,
"attempt": attempt + 1,
"max_retries": max_retries,
"amount": amount,
"unit": unit,
"mint": mint,
},
)
else:
except Exception as error:
logger.error(
"Failed to create refund token after all retries",
"Failed to create refund token",
extra={
"error": str(e),
"error_type": type(e).__name__,
"attempt": attempt + 1,
"max_retries": max_retries,
"error": str(error),
"error_type": type(error).__name__,
"amount": amount,
"unit": unit,
"mint": mint,
},
)
if refund_token is None:
raise HTTPException(
status_code=401,
detail={
"error": {
"message": f"failed to create refund after {max_retries} attempts: {str(last_exception)}",
"message": f"failed to create refund: {error}",
"type": "invalid_request_error",
"code": "send_token_failed",
}
},
)
) from error
logger.info(
"Refund token created successfully",
@@ -3683,7 +3723,6 @@ class BaseUpstreamProvider:
"amount": amount,
"unit": unit,
"mint": mint,
"attempt": attempt + 1,
"token_preview": refund_token[:20] + "..."
if len(refund_token) > 20
else refund_token,
@@ -5362,7 +5401,7 @@ class BaseUpstreamProvider:
except Exception as e:
logger.error(
f"Failed to refresh models cache for {self.provider_type or self.base_url}",
extra={"error": str(e), "error_type": type(e).__name__},
extra={"error": repr(e), "error_type": type(e).__name__},
)
def get_cached_models(self) -> list[Model]:
+92 -20
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,87 @@ if TYPE_CHECKING:
logger = get_logger(__name__)
_PPQ_SAFE_READ_ATTEMPTS = 3
_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0
class PPQCircuitOpenError(RuntimeError):
pass
@dataclass
class _PPQCircuitState:
consecutive_failures: int = 0
cooldown_until: float = 0.0
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
loop: asyncio.AbstractEventLoop | None = None
_ppq_circuits: dict[str, _PPQCircuitState] = {}
def _ppq_origin(url: str) -> str:
parsed = httpx.URL(url)
port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0)
return f"{parsed.scheme}://{parsed.host}:{port}"
async def _safe_read_request(
client: httpx.AsyncClient,
method: str,
url: str,
*,
headers: dict[str, str],
json: dict[str, object] | None = None,
) -> httpx.Response:
state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState())
loop = asyncio.get_running_loop()
if state.loop is not loop:
# Locks cannot be reused across event loops.
state.lock = asyncio.Lock()
state.loop = loop
async with state.lock:
remaining = state.cooldown_until - time.monotonic()
if remaining > 0:
raise PPQCircuitOpenError(
f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s"
)
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
try:
response = await client.request(method, url, headers=headers, json=json)
response.raise_for_status()
state.consecutive_failures = 0
state.cooldown_until = 0.0
return response
except (httpx.TransportError, httpx.HTTPStatusError) as error:
retryable_status = isinstance(error, httpx.HTTPStatusError) and (
error.response.status_code in {502, 503, 504}
)
if not isinstance(error, httpx.TransportError) and not retryable_status:
raise
state.consecutive_failures += 1
if attempt >= _PPQ_SAFE_READ_ATTEMPTS:
state.cooldown_until = (
time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS
)
raise
base_delay = 0.25 * (2 ** (attempt - 1))
delay = base_delay + random.uniform(0.0, base_delay)
logger.warning(
"PPQ.AI safe read failed; retrying",
extra={
"url": url,
"attempt": attempt,
"max_attempts": _PPQ_SAFE_READ_ATTEMPTS,
"backoff_seconds": round(delay, 3),
"error": repr(error),
"error_type": type(error).__name__,
},
)
await asyncio.sleep(delay)
raise RuntimeError("unreachable")
class PPQAIModelPricing(BaseModel):
ui: Optional[dict[str, float]] = None
@@ -123,10 +208,8 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
url = f"{self.base_url}/models"
headers = {"Authorization": f"Bearer {self.api_key}"}
try:
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url, headers=headers)
response.raise_for_status()
response = await _safe_read_request(client, "GET", url, headers=headers)
data = response.json()
models_data = data.get("data", [])
@@ -157,9 +240,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
if or_model:
input_price = None
if ppqai_model.pricing.api:
input_price = ppqai_model.pricing.api.get(
"input_per_1M"
)
input_price = ppqai_model.pricing.api.get("input_per_1M")
elif ppqai_model.pricing.input_per_1M_tokens:
input_price = ppqai_model.pricing.input_per_1M_tokens
@@ -168,9 +249,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
output_price = None
if ppqai_model.pricing.api:
output_price = ppqai_model.pricing.api.get(
"output_per_1M"
)
output_price = ppqai_model.pricing.api.get("output_per_1M")
elif ppqai_model.pricing.output_per_1M_tokens:
output_price = ppqai_model.pricing.output_per_1M_tokens
@@ -233,13 +312,6 @@ 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 []
async def on_upstream_error_redirect(
self, status_code: int, error_message: str
) -> None:
@@ -360,8 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
)
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(url, headers=headers)
response.raise_for_status()
response = await _safe_read_request(client, "GET", url, headers=headers)
status_data = response.json()
is_paid = status_data.get("status") == "Settled"
@@ -460,8 +531,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
logger.debug("Checking PPQ.AI account balance", extra={"url": url})
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.post(url, headers=headers, json={})
response.raise_for_status()
response = await _safe_read_request(
client, "POST", url, headers=headers, json={}
)
balance_data = response.json()
logger.debug(
+99 -20
View File
@@ -14,6 +14,7 @@ from typing import AsyncGenerator, TypedDict
import httpx
from cashu.core.base import MeltQuote, MeltQuoteState, MintQuote, Proof, Token
from cashu.core.mint_info import MintInfo as _CashuMintInfo
from cashu.wallet.crud import get_keysets as get_cashu_keysets
from cashu.wallet.helpers import deserialize_token_from_string
from cashu.wallet.wallet import Wallet as _CashuWallet
from pydantic_core import PydanticUndefined
@@ -121,6 +122,12 @@ def _mints_to_inspect() -> list[str]:
return mint_urls
_WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS = 30
_WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS = 300
_mint_metadata_last_load: dict[str, float] = {}
_mint_metadata_load_locks: dict[str, asyncio.Lock] = {}
class Wallet(_CashuWallet):
"""Cashu adapter that preserves HTTP 429 for Routstr's mint policy."""
@@ -141,11 +148,34 @@ class Wallet(_CashuWallet):
_CashuWallet.raise_on_error_request(resp)
async def load_mint(
self, keyset_id: str = "", force_old_keysets: bool = False
self,
keyset_id: str = "",
force_old_keysets: bool = False,
*,
force_refresh: bool = False,
) -> None:
mint_url = str(self.url)
lock = _mint_metadata_load_locks.setdefault(mint_url, asyncio.Lock())
async with lock:
now = time.monotonic()
last = _mint_metadata_last_load.get(mint_url)
if (
not force_refresh
and last is not None
and now - last < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
):
try:
await self.load_keysets_from_db()
await self.activate_keyset(keyset_id)
await self.load_mint_info(reload=False)
return
except Exception:
pass
await self.load_mint_keysets(force_old_keysets)
await self.activate_keyset(keyset_id)
await self.load_mint_info(reload=True)
_mint_metadata_last_load[mint_url] = time.monotonic()
class MintConnectionError(Exception):
@@ -491,7 +521,11 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
async def _send_locked(
amount: int, unit: str, mint_url: str | None = None
amount: int,
unit: str,
mint_url: str | None = None,
*,
owner_only: bool = False,
) -> tuple[int, str]:
effective_mint_url = await find_trusted_mint_with_funds(
amount, unit, mint_url, force_reload=True
@@ -501,6 +535,12 @@ async def _send_locked(
wallet, effective_mint_url, unit, not_reserved=True
)
proofs_for_mint = sum(proof.amount for proof in proofs)
if owner_only:
owner_balance = await _owner_balance_for_mint_and_unit(
effective_mint_url, unit, proofs_for_mint
)
if owner_balance < amount:
raise ValueError("Owner Cashu balance is insufficient for auto top-up")
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
@@ -547,6 +587,13 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str
return token
async def send_token_from_owner_locked(
amount: int, unit: str, mint_url: str | None = None
) -> str:
_, token = await _send_locked(amount, unit, mint_url, owner_only=True)
return token
class Bolt11PaymentNotAttempted(Exception):
"""The invoice was definitively not paid, so the attempt can be retried.
@@ -1048,6 +1095,7 @@ async def _request_mint_with_fallback(
mint_url,
settings.primary_mint_unit,
retry_on_rate_limit=False,
load_proofs=False,
)
quote = await run_mint_operation(
lambda: wallet.request_mint(amount),
@@ -1179,6 +1227,7 @@ async def _calculate_swap_amount(
lambda: token_wallet.melt_quote(dummy_mint_quote.request),
op_name="swap_fee_est_melt_quote",
mint_url=token_mint_url,
retry_timeouts=False,
)
fee_reserve = dummy_melt_quote.fee_reserve
@@ -1416,6 +1465,7 @@ async def swap_to_trusted_mint(
lambda: token_wallet.melt_quote(mint_quote.request),
op_name="swap_melt_quote",
mint_url=token_obj.mint,
retry_timeouts=False,
)
except Exception as error:
if is_mint_connection_error(error):
@@ -1788,21 +1838,34 @@ async def _credit_balance_locked(
)
return amount
except Exception as e:
logger.error(
"credit_balance: Error during token redemption",
extra={"error": str(e), "error_type": type(e).__name__},
classification = classify_redemption_error(e)
expected_codes = {
"cashu_token_already_spent",
"cashu_source_mint_unreachable",
"cashu_mint_unreachable",
"cashu_mint_rate_limited",
}
log = (
logger.info
if classification is not None and classification[3] in expected_codes
else logger.error
)
log(
"credit_balance: Token redemption failed",
extra={
"error": str(e),
"error_type": type(e).__name__,
"error_code": classification[3] if classification else None,
},
)
raise
_wallets: dict[str, Wallet] = {}
# Proofs require a shorter refresh interval than remote mint metadata.
_wallet_last_load: dict[str, float] = {}
_wallet_last_mint_load: dict[str, float] = {}
_wallet_load_locks: dict[str, asyncio.Lock] = {}
# Minimum seconds between full mint info + proof reloads for the same
# wallet. Prevents redundant mint API calls when get_wallet(load=True)
# is called rapidly by multiple background tasks (balance fetch, payout,
# auto-topup all hitting get_wallet within the same cycle).
_WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30
async def get_wallet(
@@ -1811,8 +1874,9 @@ async def get_wallet(
load: bool = True,
retry_on_rate_limit: bool = True,
force_reload: bool = False,
load_proofs: bool = True,
) -> Wallet:
global _wallets, _wallet_last_load, _wallet_load_locks
global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks
id = f"{mint_url}_{unit}"
lock = _wallet_load_locks.setdefault(id, asyncio.Lock())
async with lock:
@@ -1821,18 +1885,32 @@ async def get_wallet(
if load:
now = time.monotonic()
last = _wallet_last_load.get(id)
last_mint_load = _wallet_last_mint_load.get(id)
if (
force_reload
or last is None
or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS
or last_mint_load is None
or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
):
await run_mint_operation(
lambda: _wallets[id].load_mint(),
lambda: (
_wallets[id].load_mint(force_refresh=True)
if force_reload
else _wallets[id].load_mint()
),
op_name="load_mint",
mint_url=mint_url,
retry_on_rate_limit=retry_on_rate_limit,
)
_wallet_last_mint_load[id] = time.monotonic()
if load_proofs:
last_proof_load = _wallet_last_load.get(id)
if (
force_reload
or last_proof_load is None
or now - last_proof_load
>= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS
):
await run_mint_operation(
lambda: _wallets[id].load_proofs(reload=True),
op_name="load_proofs",
@@ -1912,13 +1990,14 @@ async def _get_supported_mint_units(mint_url: str) -> list[str]:
if cached is not None and now < cached[0]:
return cached[1]
wallet = await get_wallet(mint_url, settings.primary_mint_unit, load=False)
keysets = await run_mint_operation(
lambda: wallet._get_keysets(),
op_name="get_mint_keysets",
mint_url=mint_url,
# A metadata load populates Cashu's shared keyset cache for all units.
wallet = await get_wallet(
mint_url,
settings.primary_mint_unit,
retry_on_rate_limit=False,
load_proofs=False,
)
keysets = await get_cashu_keysets(mint_url=wallet.url, db=wallet.db)
units: list[str] = []
for keyset in keysets:
if not keyset.active or keyset.unit is None:
+9 -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,14 @@ os.environ.pop("ADMIN_PASSWORD", None)
from routstr.core.db import ApiKey, get_session # noqa: E402
from routstr.core.main import app, lifespan # noqa: E402
from routstr.mint import MintRateGuard # noqa: E402
@pytest.fixture(autouse=True)
def isolate_mint_rate_guards() -> Iterator[None]:
MintRateGuard._guards.clear()
yield
MintRateGuard._guards.clear()
@pytest.fixture(scope="session")
+34 -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,13 @@ async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim(
assert audit["collected"] is True
assert "lnbc-secret-invoice" not in audit["token"]
# Reusing the deterministic claim lock must not overwrite history.
assert await _claim_ppq_topup(_row()) is None
async with create_session() as session:
state = await session.get(CashuTransaction, _ppq_state_id_for_provider(1))
assert state is not None
state.created_at = int(time.time()) - 301
session.add(state)
await session.commit()
assert await _claim_ppq_topup(_row()) is not None
async with create_session() as session:
assert await session.get(CashuTransaction, audit["id"]) is not None
+4 -17
View File
@@ -686,7 +686,7 @@ async def test_no_database_changes_during_provider_operations(
@pytest.mark.integration
@pytest.mark.asyncio
async def test_admin_routstr_topup_retries_transient_upstream_failure(
async def test_admin_routstr_topup_does_not_duplicate_invoice_on_upstream_failure(
integration_client: AsyncClient,
integration_session: Any,
) -> None:
@@ -739,16 +739,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure(
assert json["api_key"] == "sk-upstream-test"
assert headers["Authorization"] == "Bearer sk-upstream-test"
if self.calls == 1:
return MockResponse(500, {"detail": "warmup failure"})
return MockResponse(
200,
{
"bolt11": "lnbc1testinvoice",
"invoice_id": "invoice-123",
},
)
return MockResponse(500, {"detail": "ambiguous upstream failure"})
mock_client = MockAsyncClient()
@@ -759,11 +750,7 @@ async def test_admin_routstr_topup_retries_transient_upstream_failure(
json={"amount": 10},
)
assert response.status_code == 200
data = response.json()
assert data["ok"] is True
assert data["topup_data"]["payment_request"] == "lnbc1testinvoice"
assert data["topup_data"]["invoice_id"] == "invoice-123"
assert mock_client.calls == 2
assert response.status_code == 500
assert mock_client.calls == 1
finally:
admin_sessions.pop(admin_token, None)
@@ -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(
+16 -21
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,32 @@ async def test_api_key_generation_invalid_token(
async def test_duplicate_token_handling(
integration_client: AsyncClient, testmint_wallet: Any, db_snapshot: Any
) -> None:
"""Test that duplicate tokens return the same API key without double-spending"""
# Generate a valid token
amount = 500 # 500 sats
amount = 500
token = await testmint_wallet.mint_tokens(amount)
# First use of token
integration_client.headers["Authorization"] = f"Bearer {token}"
response1 = await integration_client.get("/v1/wallet/info")
assert response1.status_code == 200
response1, response2 = await asyncio.gather(
integration_client.get("/v1/wallet/info"),
integration_client.get("/v1/wallet/info"),
)
assert response1.status_code < 500
assert response2.status_code < 500
assert response1.status_code == response2.status_code == 200
api_key1 = response1.json()["api_key"]
balance1 = response1.json()["balance"]
# Capture state after first submission
await db_snapshot.capture()
# Second use of same token - should return same API key since it's already created
response2 = await integration_client.get("/v1/wallet/info")
assert response2.status_code == 200
api_key2 = response2.json()["api_key"]
balance1 = response1.json()["balance"]
balance2 = response2.json()["balance"]
# Should return the same API key and balance
assert api_key1 == api_key2
assert balance1 == balance2
assert balance1 == balance2 == amount * 1000
# Verify no additional database changes
await db_snapshot.capture()
replay = await integration_client.get("/v1/wallet/info")
assert replay.status_code == 200
assert replay.json()["api_key"] == api_key1
diff = await db_snapshot.diff()
assert len(diff["api_keys"]["added"]) == 0
assert len(diff["api_keys"]["modified"]) == 0
# Original API key should still work with original balance
integration_client.headers["Authorization"] = f"Bearer {api_key1}"
wallet_response = await integration_client.get("/v1/wallet/")
assert wallet_response.status_code == 200
+100 -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,34 @@ async def test_melt_recovery_is_findable_by_quote_after_restart(
assert all(p.melt_id == QUOTE_ID for p in found)
async def test_paid_reconciliation_invalidates_recovered_proofs_after_restart(
tmp_path: Path,
) -> None:
wallet = await _wallet(tmp_path)
await _seed_ambiguous_melt(wallet)
restarted = await _wallet(tmp_path)
remote = PostMeltQuoteResponse(
quote=QUOTE_ID,
amount=95,
unit="sat",
request="lnbc1-test",
fee_reserve=1,
state=MeltQuoteState.paid.value,
expiry=None,
payment_preimage="preimage",
)
with patch(
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
new=AsyncMock(return_value=remote),
):
reconciled = await restarted.get_melt_quote(QUOTE_ID)
assert reconciled is not None and reconciled.state == MeltQuoteState.paid
assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == []
assert await crud.get_proofs(db=restarted.db) == []
async def test_send_style_reservation_would_not_be_reconcilable(
tmp_path: Path,
) -> None:
@@ -97,19 +142,64 @@ async def test_send_style_reservation_would_not_be_reconcilable(
async def test_unpaid_reconciliation_releases_recovered_proofs_after_restart(
tmp_path: Path,
) -> None:
"""The full recovery arc: crash, restart, mint says unpaid, funds usable."""
wallet = await _wallet(tmp_path)
await _seed_ambiguous_melt(wallet)
restarted = await _wallet(tmp_path)
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
assert len(found) == 2
remote = PostMeltQuoteResponse(
quote=QUOTE_ID,
amount=95,
unit="sat",
request="lnbc1-test",
fee_reserve=1,
state=MeltQuoteState.unpaid.value,
expiry=None,
)
with patch(
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
new=AsyncMock(return_value=remote),
):
reconciled = await restarted.get_melt_quote(QUOTE_ID)
# What get_melt_quote() does on an "unpaid" answer.
await restarted.set_reserved_for_melt(found, reserved=False, quote_id=None)
released = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
assert released == []
assert reconciled is not None and reconciled.state == MeltQuoteState.unpaid
assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == []
all_proofs = await crud.get_proofs(db=restarted.db)
assert len(all_proofs) == 2
assert all(not p.reserved for p in all_proofs) # spendable again
assert all(not p.reserved for p in all_proofs)
async def test_pending_and_transport_reconciliation_keep_recovered_reservation(
tmp_path: Path,
) -> None:
wallet = await _wallet(tmp_path)
await _seed_ambiguous_melt(wallet)
restarted = await _wallet(tmp_path)
pending = PostMeltQuoteResponse(
quote=QUOTE_ID,
amount=95,
unit="sat",
request="lnbc1-test",
fee_reserve=1,
state=MeltQuoteState.pending.value,
expiry=None,
)
with patch(
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
new=AsyncMock(return_value=pending),
):
reconciled = await restarted.get_melt_quote(QUOTE_ID)
assert reconciled is not None and reconciled.state == MeltQuoteState.pending
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
assert len(found) == 2 and all(proof.reserved for proof in found)
with (
patch(
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
new=AsyncMock(side_effect=httpx.ReadTimeout("mint unavailable")),
),
pytest.raises(httpx.ReadTimeout),
):
await restarted.get_melt_quote(QUOTE_ID)
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
assert len(found) == 2 and all(proof.reserved for proof in found)
+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()
+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"},
}
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},
+31
View File
@@ -56,9 +56,40 @@ def test_non_sqlite_backend_enables_pre_ping_automatically(
assert created is fake_engine
assert factory.call_args.kwargs["pool_pre_ping"] is True
assert "timeout" not in factory.call_args.kwargs["connect_args"]
assert listen.call_count == 2
def test_file_sqlite_sets_busy_timeout_connect_arg(
monkeypatch: pytest.MonkeyPatch, tmp_path: object
) -> None:
monkeypatch.setattr(settings, "database_busy_timeout", 42.0)
fake_engine = MagicMock()
with (
patch.object(db, "create_async_engine", return_value=fake_engine) as factory,
patch.object(db.event, "listen"),
):
create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/busy.db")
assert factory.call_args.kwargs["connect_args"]["timeout"] == 42.0
def test_memory_sqlite_omits_busy_timeout_connect_arg(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "database_busy_timeout", 42.0)
fake_engine = MagicMock()
with (
patch.object(db, "create_async_engine", return_value=fake_engine) as factory,
patch.object(db.event, "listen"),
):
create_db_engine("sqlite+aiosqlite://")
assert "timeout" not in factory.call_args.kwargs["connect_args"]
@pytest.mark.asyncio
async def test_every_created_engine_warns_for_long_checkouts(
monkeypatch: pytest.MonkeyPatch, tmp_path: object
+4 -3
View File
@@ -136,19 +136,20 @@ async def test_supported_mint_units_come_from_active_keysets() -> None:
msat = MagicMock(active=False, unit="msat")
usd = MagicMock(active=True)
usd.unit.name = "usd"
wallet = MagicMock()
wallet._get_keysets = AsyncMock(return_value=[usd, msat, sat])
wallet = MagicMock(url="http://mint:3338", db=MagicMock())
get_keysets = AsyncMock(return_value=[usd, msat, sat])
with (
patch.object(settings, "primary_mint_unit", "sat"),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)),
patch("routstr.wallet.get_cashu_keysets", get_keysets),
):
units = await _get_supported_mint_units("http://mint:3338")
cached_units = await _get_supported_mint_units("http://mint:3338")
assert units == ["sat", "usd"]
assert cached_units == units
wallet._get_keysets.assert_awaited_once()
get_keysets.assert_awaited_once_with(mint_url=wallet.url, db=wallet.db)
@pytest.mark.asyncio
+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",
@@ -5,6 +5,7 @@ hand an LNURL a set of unreserved proofs, and reserving only after that check
is what stops a pre-dispatch failure from stranding proofs.
"""
import math
from collections.abc import Callable
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
@@ -27,19 +28,28 @@ LNURL_DATA = {
"max_sendable": 100_000_000,
}
# 1000 sat minus the 11 sat estimated fee reserve.
EXPECTED_QUOTE_SAT = 989
EXPECTED_QUOTE_SAT = 999
def _wallet(
quote_amount: int = EXPECTED_QUOTE_SAT,
quote_amount: int | None = None,
) -> tuple[MagicMock, list[MagicMock]]:
proofs = [MagicMock(amount=1000)]
proofs = [MagicMock(amount=1000, reserved=False)]
wallet = MagicMock(url="https://mint.test")
wallet.get_fees_for_proofs.return_value = 0
if quote_amount is None:
wallet.melt_quote = AsyncMock(
side_effect=[
MagicMock(fee_reserve=1, quote="q", amount=1000),
MagicMock(fee_reserve=1, quote="q", amount=EXPECTED_QUOTE_SAT),
]
)
else:
wallet.melt_quote = AsyncMock(
return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount)
)
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
wallet.get_fees_for_proofs = MagicMock(return_value=0)
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
wallet.set_reserved_for_send = AsyncMock()
return wallet, proofs
@@ -124,16 +134,21 @@ async def test_raw_send_to_lnurl_accepts_exact_invoice() -> None:
)
assert paid == EXPECTED_QUOTE_SAT * 1000
wallet.select_to_send.assert_awaited_once()
wallet.select_to_send.assert_not_awaited()
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=True)
wallet.melt.assert_awaited_once()
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None:
# 1_000_000 msat minus an 11 sat fee reserve leaves 989_000 msat.
wallet, proofs = _wallet(989_000)
wallet, proofs = _wallet()
proofs[0].amount = 1_000_000
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
wallet.melt_quote = AsyncMock(
side_effect=[
MagicMock(fee_reserve=1, quote="q", amount=1_000_000),
MagicMock(fee_reserve=1, quote="q", amount=999_999),
]
)
data_patch, invoice_patch = _lnurl_patches()
with data_patch, invoice_patch:
@@ -141,7 +156,46 @@ async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None:
wallet, proofs, "owner@ln.tld", "msat", amount=1_000_000
)
assert paid == 989_000
assert paid == 999_999
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion() -> (
None
):
proofs = [MagicMock(amount=1, reserved=False) for _ in range(1500)]
wallet = MagicMock(url="https://mint.test")
wallet.get_fees_for_proofs = MagicMock(
side_effect=lambda selected: math.ceil(len(selected) / 100)
)
wallet.melt_quote = AsyncMock(
side_effect=[
MagicMock(fee_reserve=10, quote="q1", amount=1500),
MagicMock(fee_reserve=10, quote="q2", amount=1475),
]
)
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
wallet.set_reserved_for_send = AsyncMock()
checkpoint = AsyncMock()
data_patch, invoice_patch = _lnurl_patches()
with data_patch, invoice_patch:
paid = await raw_send_to_lnurl(
wallet,
proofs,
"owner@ln.tld",
"sat",
amount=1500,
on_melt_quote=checkpoint,
)
assert paid == 1_475_000
assert wallet.melt_quote.await_count == 2
checkpoint.assert_awaited_once_with("q2")
wallet.select_to_send.assert_not_called()
selected = wallet.melt.await_args.kwargs["proofs"]
assert sum(proof.amount for proof in selected) == 1500
assert 1475 + 10 + wallet.get_fees_for_proofs(selected) == 1500
@pytest.mark.asyncio
@@ -291,3 +345,23 @@ async def test_send_to_lnurl_does_not_reserve_before_lnurl_validation() -> None:
assert raw_send.await_args is not None
assert raw_send.await_args.args[1] is proofs
assert raw_send.await_args.kwargs["amount"] == 1000
def test_select_melt_proofs_stops_at_minimal_cover_when_over_budget() -> None:
from routstr.payment.lnurl import _select_melt_proofs
wallet = MagicMock()
wallet.get_fees_for_proofs = MagicMock(side_effect=lambda selected: len(selected))
proofs = [MagicMock(amount=600, reserved=False) for _ in range(3)]
selected, shortfall = _select_melt_proofs(
wallet,
proofs,
quote_amount=1000,
fee_reserve=0,
gross_budget=1000,
)
assert selected is None
assert shortfall == 2
assert wallet.get_fees_for_proofs.call_count == 2
+141 -9
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
@@ -11,10 +12,19 @@ from cashu.core.base import MeltQuoteState
from routstr.core.settings import settings
from routstr.mint import MintCooldownError, MintRateGuard
from routstr.payment.lnurl import (
LNURLError,
MeltOutcomeAmbiguousError,
raw_send_to_lnurl,
)
@pytest.fixture(autouse=True)
def _clear_mint_guards() -> Iterator[None]:
MintRateGuard._guards.clear()
yield
MintRateGuard._guards.clear()
LNURL_DATA = {
"callback_url": "https://ln.tld/cb",
"min_sendable": 1_000,
@@ -22,17 +32,23 @@ LNURL_DATA = {
}
# 1000 sat minus the 11 sat estimated fee reserve.
QUOTE_AMOUNT_SAT = 989
QUOTE_AMOUNT_SAT = 999
def _wallet() -> tuple[MagicMock, list[MagicMock]]:
proofs = [MagicMock(amount=1000)]
proofs = [MagicMock(amount=1000, reserved=False)]
wallet = MagicMock(url="https://mint.test")
wallet.melt_quote = AsyncMock(
return_value=MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT)
side_effect=[
MagicMock(fee_reserve=1, quote="q", amount=1000),
MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT),
]
)
wallet.melt = AsyncMock()
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
wallet.get_fees_for_proofs = MagicMock(return_value=0)
wallet.set_reserved_for_melt = AsyncMock()
wallet.set_reserved_for_send = AsyncMock()
return wallet, proofs
@@ -50,7 +66,26 @@ def _lnurl_patches() -> tuple[Any, Any]:
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> None:
async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None:
wallet, proofs = _wallet()
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.unpaid))
wallet.get_melt_quote = AsyncMock()
data_patch, invoice_patch = _lnurl_patches()
with (
data_patch,
invoice_patch,
pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised,
):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
assert not isinstance(raised.value, MeltOutcomeAmbiguousError)
wallet.get_melt_quote.assert_not_awaited()
wallet.set_reserved_for_send.assert_any_await(proofs, reserved=False)
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_timeout_then_unpaid_remains_ambiguous() -> None:
wallet, proofs = _wallet()
async def _hang(**kwargs: object) -> None:
@@ -67,12 +102,65 @@ async def test_raw_send_to_lnurl_timeout_keeps_unpaid_outcome_ambiguous() -> Non
patch.object(settings, "mint_retry_max_attempts", 0),
data_patch,
invoice_patch,
pytest.raises(MeltOutcomeAmbiguousError, match="outcome is ambiguous"),
pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"),
):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.get_melt_quote.assert_awaited_once_with("q")
wallet.set_reserved_for_melt.assert_not_called()
assert wallet.set_reserved_for_melt.await_count == 2
wallet.set_reserved_for_melt.assert_awaited_with(
proofs, reserved=True, quote_id="q"
)
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_wrapped_transport_unpaid_remains_ambiguous() -> None:
wallet, proofs = _wallet()
async def _wrapped_transport_error(**kwargs: object) -> None:
try:
raise httpx.ReadTimeout("response lost")
except httpx.ReadTimeout as transport_error:
raise Exception("could not pay invoice") from transport_error
wallet.melt = AsyncMock(side_effect=_wrapped_transport_error)
wallet.get_melt_quote = AsyncMock(
return_value=MagicMock(state=MeltQuoteState.unpaid)
)
data_patch, invoice_patch = _lnurl_patches()
with (
patch.object(settings, "mint_retry_max_attempts", 3),
data_patch,
invoice_patch,
pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"),
):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.melt.assert_awaited_once()
wallet.get_melt_quote.assert_awaited_once_with("q")
assert wallet.set_reserved_for_melt.await_count == 2
wallet.set_reserved_for_melt.assert_awaited_with(
proofs, reserved=True, quote_id="q"
)
@pytest.mark.asyncio
async def test_raw_send_to_lnurl_does_not_retry_melt_quote_timeout() -> None:
wallet, proofs = _wallet()
wallet.melt_quote = AsyncMock(side_effect=httpx.ReadTimeout("response lost"))
data_patch, invoice_patch = _lnurl_patches()
with (
patch.object(settings, "mint_retry_max_attempts", 3),
data_patch,
invoice_patch,
pytest.raises(httpx.TimeoutException),
):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.melt_quote.assert_awaited_once()
wallet.melt.assert_not_awaited()
@pytest.mark.asyncio
@@ -98,6 +186,9 @@ async def test_raw_send_to_lnurl_timeout_reconciled_paid_is_success() -> None:
assert paid > 0
wallet.get_melt_quote.assert_awaited_once_with("q")
wallet.set_reserved_for_melt.assert_awaited_once_with(
proofs, reserved=True, quote_id="q"
)
@pytest.mark.asyncio
@@ -120,6 +211,45 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None:
wallet.get_melt_quote.assert_awaited_once_with("q")
@pytest.mark.asyncio
async def test_pending_then_immediate_unpaid_remains_reserved_and_ambiguous() -> None:
wallet, proofs = _wallet()
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending))
wallet.get_melt_quote = AsyncMock(
return_value=MagicMock(state=MeltQuoteState.unpaid)
)
data_patch, invoice_patch = _lnurl_patches()
with (
data_patch,
invoice_patch,
pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"),
):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.set_reserved_for_melt.assert_awaited_once_with(
proofs, reserved=True, quote_id="q"
)
@pytest.mark.asyncio
async def test_immediate_unpaid_reservation_failure_stays_ambiguous() -> None:
wallet, proofs = _wallet()
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending))
wallet.get_melt_quote = AsyncMock(
return_value=MagicMock(state=MeltQuoteState.unpaid)
)
wallet.set_reserved_for_melt = AsyncMock(side_effect=OSError("db locked"))
data_patch, invoice_patch = _lnurl_patches()
with (
data_patch,
invoice_patch,
pytest.raises(MeltOutcomeAmbiguousError, match="could not be restored"),
):
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
@pytest.mark.asyncio
@pytest.mark.parametrize("rate_error", ["cooldown", "http_429"])
async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs(
@@ -153,7 +283,8 @@ async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs(
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.melt.assert_not_awaited()
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=False)
assert wallet.set_reserved_for_send.await_count == 2
wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False)
@pytest.mark.asyncio
@@ -178,7 +309,8 @@ async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None:
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
wallet.melt.assert_awaited_once()
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=False)
assert wallet.set_reserved_for_send.await_count == 2
wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False)
MintRateGuard._guards.pop(str(wallet.url), None)
+77
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,82 @@ async def test_cashu_429_dispatches_through_wallet_override() -> None:
await wallet.mint_quote(1, Unit.sat)
@pytest.mark.asyncio
async def test_wrapped_transport_failure_opens_central_cooldown() -> None:
mint_url = "https://transport-failure.test"
MintRateGuard._guards.pop(mint_url, None)
async def wrapped_failure() -> None:
try:
raise httpx.ReadTimeout("body stalled")
except httpx.ReadTimeout as error:
raise Exception("wallet wrapper") from error
with pytest.raises(Exception, match="wallet wrapper"):
await run_mint_operation(
wrapped_failure,
mint_url=mint_url,
retry_timeouts=False,
)
guard = MintRateGuard.get(mint_url)
assert guard.cooldown_remaining() > 29
probe = AsyncMock()
async with fail_fast_mint_operations():
with pytest.raises(MintCooldownError):
await guard.run(probe)
probe.assert_not_awaited()
MintRateGuard._guards.pop(mint_url, None)
@pytest.mark.asyncio
async def test_timeout_retry_succeeds_without_opening_cooldown() -> None:
from routstr.core.settings import settings
mint_url = "https://retryable-timeout.test"
MintRateGuard._guards.pop(mint_url, None)
calls = 0
async def flaky() -> str:
nonlocal calls
calls += 1
if calls == 1:
raise httpx.ReadTimeout("first attempt stalled")
return "ok"
with (
patch.object(settings, "mint_retry_max_attempts", 2),
patch("routstr.mint.asyncio.sleep", AsyncMock()),
):
result = await run_mint_operation(flaky, mint_url=mint_url)
assert result == "ok"
assert calls == 2
assert MintRateGuard.get(mint_url).cooldown_remaining() == 0.0
MintRateGuard._guards.pop(mint_url, None)
@pytest.mark.asyncio
async def test_exhausted_timeout_retries_open_transport_cooldown() -> None:
from routstr.core.settings import settings
mint_url = "https://exhausted-timeout.test"
MintRateGuard._guards.pop(mint_url, None)
async def always_timeout() -> None:
raise httpx.ReadTimeout("stalled")
with (
patch.object(settings, "mint_retry_max_attempts", 1),
patch("routstr.mint.asyncio.sleep", AsyncMock()),
pytest.raises(httpx.TimeoutException),
):
await run_mint_operation(always_timeout, mint_url=mint_url)
assert MintRateGuard.get(mint_url).cooldown_remaining() > 29
MintRateGuard._guards.pop(mint_url, None)
async def test_guard_concurrency_change_preserves_cooldown_state() -> None:
from routstr.core.settings import settings
+19 -1
View File
@@ -1,7 +1,8 @@
"""Persisted mint preferences must not bypass the configured trusted set."""
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from routstr.core.settings import settings
@@ -31,6 +32,23 @@ async def test_untrusted_allowed_mints_fall_back_to_trusted_set() -> None:
assert attempted == [TRUSTED]
async def test_mint_quote_timeout_is_not_retried() -> None:
wallet = MagicMock()
wallet.request_mint = AsyncMock(side_effect=httpx.ReadTimeout("response lost"))
with (
patch.object(settings, "primary_mint", TRUSTED),
patch.object(settings, "cashu_mints", [TRUSTED]),
patch.object(settings, "mint_retry_max_attempts", 3),
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
patch("routstr.lightning.mint_cooldown_remaining", return_value=0.0),
pytest.raises(Exception),
):
await _request_mint_with_fallback(10)
wallet.request_mint.assert_awaited_once_with(10)
async def test_trusted_allowed_mints_are_used_verbatim() -> None:
attempted: list[str] = []
+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()
+47
View File
@@ -0,0 +1,47 @@
from unittest.mock import AsyncMock, patch
import httpx
import pytest
from fastapi import HTTPException
from routstr.upstream.base import BaseUpstreamProvider
@pytest.mark.asyncio
async def test_send_refund_does_not_retry_ambiguous_token_creation() -> None:
provider = object.__new__(BaseUpstreamProvider)
send_token = AsyncMock(side_effect=httpx.ReadTimeout("swap response lost"))
store = AsyncMock()
with (
patch("routstr.upstream.base.send_token", send_token),
patch("routstr.upstream.base.store_cashu_transaction", store),
pytest.raises(HTTPException) as raised,
):
await provider.send_refund(10, "sat", mint="https://mint.test")
assert raised.value.status_code == 401
send_token.assert_awaited_once_with(
10, unit="sat", mint_url="https://mint.test"
)
store.assert_not_awaited()
@pytest.mark.asyncio
async def test_send_refund_does_not_retry_or_store_on_generic_failure() -> None:
provider = object.__new__(BaseUpstreamProvider)
send_token = AsyncMock(side_effect=Exception("mint rejected swap"))
store = AsyncMock()
with (
patch("routstr.upstream.base.send_token", send_token),
patch("routstr.upstream.base.store_cashu_transaction", store),
pytest.raises(HTTPException) as raised,
):
await provider.send_refund(10, "sat", mint="https://mint.test")
assert raised.value.status_code == 401
send_token.assert_awaited_once_with(
10, unit="sat", mint_url="https://mint.test"
)
store.assert_not_awaited()
@@ -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
+126
View File
@@ -26,6 +26,7 @@ from routstr.wallet import (
recieve_token,
send,
send_token,
send_token_from_owner_locked,
)
@@ -40,13 +41,19 @@ def isolate_wallet_runtime_state() -> Generator[None, None, None]:
wallet_module._MintRateGuard._guards.clear()
wallet_module._wallets.clear()
wallet_module._wallet_last_load.clear()
wallet_module._wallet_last_mint_load.clear()
wallet_module._wallet_load_locks.clear()
wallet_module._mint_metadata_last_load.clear()
wallet_module._mint_metadata_load_locks.clear()
yield
settings.mint_max_concurrency = original_concurrency
wallet_module._MintRateGuard._guards.clear()
wallet_module._wallets.clear()
wallet_module._wallet_last_load.clear()
wallet_module._wallet_last_mint_load.clear()
wallet_module._wallet_load_locks.clear()
wallet_module._mint_metadata_last_load.clear()
wallet_module._mint_metadata_load_locks.clear()
@pytest.mark.asyncio
@@ -66,6 +73,63 @@ async def test_get_balance() -> None:
assert balance == 50000
@pytest.mark.asyncio
async def test_wallet_metadata_is_reused_across_units() -> None:
from routstr.wallet import Wallet
sat_wallet = MagicMock(url="http://mint:3338")
sat_wallet.load_mint_keysets = AsyncMock()
sat_wallet.activate_keyset = AsyncMock()
sat_wallet.load_mint_info = AsyncMock()
sat_wallet.load_keysets_from_db = AsyncMock()
msat_wallet = MagicMock(url="http://mint:3338")
msat_wallet.load_mint_keysets = AsyncMock()
msat_wallet.activate_keyset = AsyncMock()
msat_wallet.load_mint_info = AsyncMock()
msat_wallet.load_keysets_from_db = AsyncMock()
with patch("routstr.wallet.time.monotonic", return_value=1000.0):
await Wallet.load_mint(sat_wallet)
await Wallet.load_mint(msat_wallet)
sat_wallet.load_mint_keysets.assert_awaited_once_with(False)
sat_wallet.load_mint_info.assert_awaited_once_with(reload=True)
msat_wallet.load_mint_keysets.assert_not_awaited()
msat_wallet.load_keysets_from_db.assert_awaited_once_with()
msat_wallet.load_mint_info.assert_awaited_once_with(reload=False)
@pytest.mark.asyncio
async def test_get_wallet_refreshes_local_proofs_without_reloading_mint() -> None:
from routstr import wallet as wallet_module
from routstr.wallet import get_wallet
mock_wallet = Mock(load_mint=AsyncMock(), load_proofs=AsyncMock())
with (
patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)),
patch("routstr.wallet.time.monotonic", return_value=1000.0),
):
await get_wallet("http://mint:3338", "sat")
wallet_module._wallet_last_load["http://mint:3338_sat"] = 900.0
await get_wallet("http://mint:3338", "sat")
assert mock_wallet.load_mint.await_count == 1
assert mock_wallet.load_proofs.await_count == 2
@pytest.mark.asyncio
async def test_get_wallet_quote_only_skips_proof_reload() -> None:
from routstr.wallet import get_wallet
mock_wallet = Mock(load_mint=AsyncMock(), load_proofs=AsyncMock())
with patch("routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)):
await get_wallet("http://mint:3338", "sat", load_proofs=False)
mock_wallet.load_mint.assert_awaited_once_with()
mock_wallet.load_proofs.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_wallet_force_reload_bypasses_reload_interval() -> None:
from routstr.wallet import get_wallet
@@ -392,6 +456,33 @@ async def test_send_token() -> None:
assert token == "test_token"
@pytest.mark.asyncio
async def test_owner_only_token_rejects_customer_backed_proofs() -> None:
mint = "http://mint:3338"
proof = Mock(amount=1000, reserved=False)
wallet = Mock(keysets={}, proofs=[proof], select_to_send=AsyncMock())
with (
patch(
"routstr.wallet.find_trusted_mint_with_funds",
AsyncMock(return_value=mint),
),
patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)),
patch(
"routstr.wallet.get_proofs_per_mint_and_unit",
return_value=[proof],
),
patch(
"routstr.wallet._owner_balance_for_mint_and_unit",
AsyncMock(return_value=50),
),
pytest.raises(ValueError, match="Owner Cashu balance"),
):
await send_token_from_owner_locked(100, "sat", mint)
wallet.select_to_send.assert_not_awaited()
@pytest.mark.asyncio
async def test_release_token_reservation_unreserves_local_proofs() -> None:
from routstr.wallet import release_token_reservation
@@ -632,6 +723,38 @@ async def test_credit_balance() -> None:
assert mock_session.refresh.called
@pytest.mark.asyncio
async def test_concurrent_duplicate_token_credits_exactly_once() -> None:
key = Mock(balance=0, hashed_key="duplicate-key")
session = AsyncMock()
session.exec.return_value.rowcount = 1
session.refresh = AsyncMock()
receive = AsyncMock(
side_effect=[
(1000, "sat", "https://mint.test"),
ValueError("Mint Error: proofs already spent (Code: 11001)"),
]
)
store = AsyncMock()
with (
patch("routstr.wallet.recieve_token", receive),
patch("routstr.wallet.store_cashu_transaction", store),
):
results = await asyncio.gather(
credit_balance("cashuAduplicate", key, session),
credit_balance("cashuAduplicate", key, session),
return_exceptions=True,
)
assert sum(result == 1_000_000 for result in results) == 1
failure = next(result for result in results if isinstance(result, Exception))
classified = classify_redemption_error(failure)
assert classified is not None and classified[3] == "cashu_token_already_spent"
assert session.exec.await_count == 1
store.assert_awaited_once()
@pytest.mark.asyncio
async def test_credit_balance_constrains_redemption_to_key_mint() -> None:
key_mint = "http://key-mint:3338"
@@ -2970,6 +3093,7 @@ async def test_load_mint_propagates_rate_limit() -> None:
from routstr.wallet import Wallet
wallet = Wallet.__new__(Wallet)
wallet.url = "https://rate-limited-mint.example"
error = MintRateLimitedError(
"Cashu mint rate limited",
request=httpx.Request("GET", "https://mint.example/v1/keysets"),
@@ -2987,6 +3111,7 @@ async def test_load_mint_propagates_connection_error() -> None:
from routstr.wallet import Wallet
wallet = Wallet.__new__(Wallet)
wallet.url = "https://unavailable-mint.example"
error = httpx.ConnectError("mint unavailable")
with (
patch.object(wallet, "load_mint_keysets", new=AsyncMock(side_effect=error)),
@@ -3002,6 +3127,7 @@ async def test_load_mint_runs_keysets_activation_and_info() -> None:
from routstr.wallet import Wallet
wallet = Wallet.__new__(Wallet)
wallet.url = "https://mint-load.example"
with (
patch.object(wallet, "load_mint_keysets", new=AsyncMock()) as load_keysets,
patch.object(wallet, "activate_keyset", new=AsyncMock()) as activate,