mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge pull request #693 from Routstr/fix/cashu-mint-interactions
fix: harden wallet and Cashu operations
This commit is contained in:
@@ -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
@@ -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,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
+27
-41
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import secrets
|
||||
@@ -1415,50 +1414,37 @@ async def initiate_provider_topup(
|
||||
else {}
|
||||
)
|
||||
|
||||
last_status_code = 500
|
||||
last_error_detail: object = "Failed to create top-up invoice"
|
||||
# Quote creation is unsafe to retry without idempotency.
|
||||
resp = await client.post(
|
||||
f"{clean_url}/v1/balance/lightning/invoice",
|
||||
json=request_json,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
# Some upstream Routstr nodes fail the first invoice request after warm-up
|
||||
# and succeed immediately on retry. Retry once here so the UI stays single-click.
|
||||
for attempt in range(2):
|
||||
resp = await client.post(
|
||||
f"{clean_url}/v1/balance/lightning/invoice",
|
||||
json=request_json,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
return {
|
||||
"ok": True,
|
||||
"topup_data": {
|
||||
"payment_request": data.get("bolt11"),
|
||||
"invoice_id": data.get("invoice_id"),
|
||||
"status": "pending",
|
||||
},
|
||||
}
|
||||
|
||||
logger.error(
|
||||
f"Upstream topup request failed: {resp.text}",
|
||||
extra={
|
||||
"provider_id": provider_id,
|
||||
"attempt": attempt + 1,
|
||||
"status_code": resp.status_code,
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
return {
|
||||
"ok": True,
|
||||
"topup_data": {
|
||||
"payment_request": data.get("bolt11"),
|
||||
"invoice_id": data.get("invoice_id"),
|
||||
"status": "pending",
|
||||
},
|
||||
)
|
||||
try:
|
||||
last_error_detail = resp.json()
|
||||
except Exception:
|
||||
last_error_detail = resp.text
|
||||
last_status_code = resp.status_code
|
||||
|
||||
if resp.status_code < 500 or attempt == 1:
|
||||
break
|
||||
|
||||
await asyncio.sleep(0.2)
|
||||
}
|
||||
|
||||
logger.error(
|
||||
f"Upstream topup request failed: {resp.text}",
|
||||
extra={
|
||||
"provider_id": provider_id,
|
||||
"status_code": resp.status_code,
|
||||
},
|
||||
)
|
||||
try:
|
||||
error_detail: object = resp.json()
|
||||
except Exception:
|
||||
error_detail = resp.text
|
||||
raise HTTPException(
|
||||
status_code=last_status_code, detail=last_error_detail
|
||||
status_code=resp.status_code, detail=error_detail
|
||||
)
|
||||
|
||||
upstream_instance = _instantiate_provider(provider)
|
||||
|
||||
+7
-1
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
+104
-28
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import math
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, TypedDict
|
||||
|
||||
@@ -10,8 +9,8 @@ from cashu.core.base import MeltQuoteState
|
||||
from cashu.wallet.wallet import Proof, Wallet
|
||||
|
||||
from ..mint import (
|
||||
MINT_TRANSPORT_EXCEPTIONS,
|
||||
is_mint_rate_limited,
|
||||
is_mint_transport_error,
|
||||
run_mint_operation,
|
||||
)
|
||||
|
||||
@@ -226,6 +225,38 @@ async def get_lnurl_invoice(
|
||||
return invoice_data["pr"], invoice_data
|
||||
|
||||
|
||||
def _select_melt_proofs(
|
||||
wallet: Wallet,
|
||||
proofs: list[Proof],
|
||||
*,
|
||||
quote_amount: int,
|
||||
fee_reserve: int,
|
||||
gross_budget: int,
|
||||
) -> tuple[list[Proof] | None, int]:
|
||||
"""Select proofs that cover the quote and exact NUT-02 input fees.
|
||||
|
||||
Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a
|
||||
wallet's full balance. Melts accept overpayment and return change, so a
|
||||
bounded, largest-first selection is both safer and minimizes input fees.
|
||||
"""
|
||||
selected: list[Proof] = []
|
||||
selected_amount = 0
|
||||
required = quote_amount + fee_reserve
|
||||
for proof in sorted(proofs, key=lambda item: item.amount, reverse=True):
|
||||
if getattr(proof, "reserved", False) is True:
|
||||
continue
|
||||
selected.append(proof)
|
||||
selected_amount += proof.amount
|
||||
input_fees = int(wallet.get_fees_for_proofs(selected))
|
||||
required = quote_amount + fee_reserve + input_fees
|
||||
if selected_amount >= required:
|
||||
if required <= gross_budget:
|
||||
return selected, 0
|
||||
# Covered but over budget; more proofs only raise input fees.
|
||||
break
|
||||
return None, max(1, required - min(selected_amount, gross_budget))
|
||||
|
||||
|
||||
async def raw_send_to_lnurl(
|
||||
wallet: Wallet,
|
||||
proofs: list[Proof],
|
||||
@@ -281,36 +312,51 @@ async def raw_send_to_lnurl(
|
||||
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
|
||||
)
|
||||
|
||||
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1
|
||||
estimated_fees_msat = estimated_fees_sat * 1000
|
||||
final_amount = amount_msat - estimated_fees_msat
|
||||
final_amount = amount_msat
|
||||
|
||||
bolt11_invoice, _ = await get_lnurl_invoice(
|
||||
lnurl_data["callback_url"], final_amount
|
||||
)
|
||||
|
||||
melt_quote_resp = await run_mint_operation(
|
||||
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
||||
op_name="lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
)
|
||||
|
||||
# The invoice comes from the LNURL service, so its amount is untrusted. The
|
||||
# melt quote is the mint's own reading of it, and it must match what we
|
||||
# asked to send. Checked before the checkpoint and before reserving, so a
|
||||
# mismatch leaves no durable state and no locked proofs behind.
|
||||
quoted_amount = int(melt_quote_resp.amount)
|
||||
expected_amount = final_amount // 1000 if unit == "sat" else final_amount
|
||||
if quoted_amount != expected_amount:
|
||||
raise LNURLError(
|
||||
f"LNURL invoice amount does not match the requested amount "
|
||||
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
|
||||
selected_proofs: list[Proof] | None = None
|
||||
# Find the largest amount covered by the budget after reserve and input fees.
|
||||
for _ in range(8):
|
||||
if final_amount < lnurl_data["min_sendable"]:
|
||||
raise LNURLError("Cashu melt fees leave no payable LNURL amount")
|
||||
bolt11_invoice, _ = await get_lnurl_invoice(
|
||||
lnurl_data["callback_url"], final_amount
|
||||
)
|
||||
melt_quote_resp = await run_mint_operation(
|
||||
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
||||
op_name="lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
# Quote creation is unsafe to retry without idempotency.
|
||||
retry_timeouts=False,
|
||||
)
|
||||
|
||||
quoted_amount = int(melt_quote_resp.amount)
|
||||
expected_amount = final_amount // 1000 if unit == "sat" else final_amount
|
||||
if quoted_amount != expected_amount:
|
||||
raise LNURLError(
|
||||
f"LNURL invoice amount does not match the requested amount "
|
||||
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
|
||||
)
|
||||
|
||||
selected_proofs, shortfall = _select_melt_proofs(
|
||||
wallet,
|
||||
proofs,
|
||||
quote_amount=quoted_amount,
|
||||
fee_reserve=int(melt_quote_resp.fee_reserve),
|
||||
gross_budget=amount,
|
||||
)
|
||||
if selected_proofs is not None:
|
||||
break
|
||||
final_amount -= shortfall * (1000 if unit == "sat" else 1)
|
||||
else:
|
||||
raise LNURLError("Cashu melt fees exceed the requested gross amount")
|
||||
|
||||
if on_melt_quote is not None:
|
||||
await on_melt_quote(melt_quote_resp.quote)
|
||||
|
||||
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
|
||||
assert selected_proofs is not None
|
||||
proofs = selected_proofs
|
||||
await wallet.set_reserved_for_send(proofs, reserved=True)
|
||||
|
||||
try:
|
||||
melt_response = await run_mint_operation(
|
||||
@@ -331,15 +377,29 @@ async def raw_send_to_lnurl(
|
||||
# reserved as though a Lightning payment could still settle.
|
||||
await wallet.set_reserved_for_send(proofs, reserved=False)
|
||||
raise
|
||||
if not isinstance(error, MINT_TRANSPORT_EXCEPTIONS):
|
||||
if not is_mint_transport_error(error):
|
||||
raise
|
||||
# Cashu clears reservations on transport errors despite an unknown outcome.
|
||||
try:
|
||||
await wallet.set_reserved_for_melt(
|
||||
proofs, reserved=True, quote_id=melt_quote_resp.quote
|
||||
)
|
||||
except Exception as reservation_error:
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous and its proof reservation could not "
|
||||
"be restored; proofs must not be retried"
|
||||
) from reservation_error
|
||||
melt_response = None
|
||||
melt_error: BaseException | None = error
|
||||
else:
|
||||
melt_error = None
|
||||
|
||||
if getattr(melt_response, "state", None) == MeltQuoteState.paid:
|
||||
melt_state = getattr(melt_response, "state", None)
|
||||
if melt_state == MeltQuoteState.paid:
|
||||
return final_amount
|
||||
if melt_state == MeltQuoteState.unpaid:
|
||||
await wallet.set_reserved_for_send(proofs, reserved=False)
|
||||
raise LNURLError("Cashu mint confirmed that the melt was unpaid")
|
||||
|
||||
try:
|
||||
quote = await run_mint_operation(
|
||||
@@ -347,6 +407,8 @@ async def raw_send_to_lnurl(
|
||||
op_name="reconcile_lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
retry_timeouts=False,
|
||||
# Reconciliation must bypass the cooldown opened by this failure.
|
||||
allow_during_cooldown=True,
|
||||
)
|
||||
except Exception as reconciliation_error:
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
@@ -356,6 +418,20 @@ async def raw_send_to_lnurl(
|
||||
|
||||
if quote is not None and quote.state == MeltQuoteState.paid:
|
||||
return final_amount
|
||||
if quote is not None and quote.state == MeltQuoteState.unpaid:
|
||||
# A just-dispatched quote can briefly report unpaid before transitioning.
|
||||
try:
|
||||
await wallet.set_reserved_for_melt(
|
||||
proofs, reserved=True, quote_id=melt_quote_resp.quote
|
||||
)
|
||||
except Exception as reservation_error:
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous and its proof reservation could not "
|
||||
"be restored; proofs must not be retried"
|
||||
) from reservation_error
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous; an immediate unpaid state is not final"
|
||||
) from melt_error
|
||||
|
||||
state = getattr(getattr(quote, "state", None), "value", "unknown")
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
|
||||
+153
-81
@@ -15,9 +15,6 @@ from ..core.db import (
|
||||
UpstreamProviderRow,
|
||||
create_session,
|
||||
)
|
||||
from ..core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from ..payment.price import sats_usd_price
|
||||
from ..wallet import (
|
||||
Bolt11PaymentAmbiguous,
|
||||
@@ -27,7 +24,7 @@ from ..wallet import (
|
||||
maximum_owner_cashu_balance_sats,
|
||||
prepare_bolt11_payment,
|
||||
release_token_reservation,
|
||||
send_token,
|
||||
send_token_from_owner_locked,
|
||||
token_mint_url,
|
||||
wallet_operation_guard,
|
||||
)
|
||||
@@ -49,6 +46,7 @@ PPQ_PHASES = frozenset({PPQ_PHASE_CLAIMED, PPQ_PHASE_IN_FLIGHT, PPQ_PHASE_RECONC
|
||||
PPQ_SETTLEMENT_ATTEMPTS = 5
|
||||
PPQ_SETTLEMENT_POLL_SECONDS = 2
|
||||
PPQ_PENDING_TTL_SECONDS = 15 * 60
|
||||
PPQ_SETTLED_COOLDOWN_SECONDS = 5 * 60
|
||||
PPQ_MAX_INVOICE_PREMIUM = 1.10
|
||||
PPQ_MIN_TOPUP_USD = 1
|
||||
PPQ_MAX_TOPUP_USD = 500
|
||||
@@ -99,7 +97,7 @@ async def periodic_auto_topup() -> None:
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Auto top-up cycle failed",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
extra={"error": repr(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS)
|
||||
@@ -130,7 +128,7 @@ async def _run_auto_topup_cycle() -> None:
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"base_url": row.base_url,
|
||||
"error": str(e),
|
||||
"error": repr(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
@@ -161,7 +159,11 @@ async def _reconcile_all_ppq_claims() -> set[int]:
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"PPQ claim reconciliation failed",
|
||||
extra={"provider_id": row.id, "error": str(e)},
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"error": repr(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
return active_provider_ids
|
||||
|
||||
@@ -363,67 +365,59 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
)
|
||||
|
||||
try:
|
||||
token = await send_token(amount, "sat", mint_url)
|
||||
async with wallet_operation_guard():
|
||||
# Keep the spend cap and audit mutation in one wallet lock.
|
||||
spent_24h_sats = await _routstr_spent_last_24h_sats()
|
||||
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
|
||||
raise ValueError("Routstr auto top-up daily spend cap reached")
|
||||
token = await send_token_from_owner_locked(amount, "sat", mint_url)
|
||||
actual_mint_url = token_mint_url(token, mint_url)
|
||||
try:
|
||||
await _persist_routstr_token_and_mark_sent(
|
||||
row,
|
||||
operation_id,
|
||||
expected_sats=expected_sats,
|
||||
token=token,
|
||||
amount=amount,
|
||||
mint_url=actual_mint_url,
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Aborting auto top-up because its token and sent claim "
|
||||
"could not be persisted atomically",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
try:
|
||||
await release_token_reservation(token)
|
||||
except Exception as error:
|
||||
logger.critical(
|
||||
"Failed to release untracked auto-topup token",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"mint_url": actual_mint_url,
|
||||
"error": repr(error),
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Auto-topup token was released after persistence failed",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to create cashu token for auto top-up",
|
||||
logger.warning(
|
||||
"Failed to create or persist cashu token for auto top-up",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"amount": amount,
|
||||
"mint_url": mint_url,
|
||||
"error": str(e),
|
||||
"error": repr(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
await _release_routstr_claim(row, operation_id)
|
||||
return
|
||||
|
||||
actual_mint_url = token_mint_url(token, mint_url)
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=actual_mint_url,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Aborting auto top-up because its cashu token could not be persisted",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
try:
|
||||
await release_token_reservation(token)
|
||||
except Exception as error:
|
||||
logger.critical(
|
||||
"Failed to release untracked auto-topup token",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"mint_url": actual_mint_url,
|
||||
"error": str(error),
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Auto-topup token was released after persistence failed",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
await _release_routstr_claim(row, operation_id)
|
||||
return
|
||||
|
||||
# Move the claim before the network call, not after: a worker that dies
|
||||
# mid-request must leave behind a claim that says a token may already be
|
||||
# with the peer.
|
||||
await _mark_routstr_sent(
|
||||
row,
|
||||
operation_id,
|
||||
expected_sats=expected_sats,
|
||||
token=token,
|
||||
amount=amount,
|
||||
mint_url=actual_mint_url,
|
||||
)
|
||||
|
||||
result = await provider.topup(token)
|
||||
|
||||
if "error" in result:
|
||||
@@ -705,7 +699,7 @@ async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) ->
|
||||
)
|
||||
|
||||
|
||||
async def _mark_routstr_sent(
|
||||
async def _persist_routstr_token_and_mark_sent(
|
||||
row: UpstreamProviderRow,
|
||||
operation_id: str,
|
||||
*,
|
||||
@@ -714,20 +708,59 @@ async def _mark_routstr_sent(
|
||||
amount: int,
|
||||
mint_url: str,
|
||||
) -> None:
|
||||
claim = await _current_routstr_claim(row)
|
||||
failures = claim.failures if claim else 0
|
||||
if not await _advance_routstr_claim(
|
||||
row,
|
||||
operation_id,
|
||||
deadline=int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS,
|
||||
phase=ROUTSTR_PHASE_SENT,
|
||||
expected_sats=expected_sats,
|
||||
failures=failures,
|
||||
token=token,
|
||||
amount=amount,
|
||||
mint_url=mint_url,
|
||||
):
|
||||
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||
state_id = _routstr_state_id(row)
|
||||
async with create_session() as session:
|
||||
state = await session.get(CashuTransaction, state_id)
|
||||
claim = _parse_routstr_request_id(state.request_id if state else None)
|
||||
if (
|
||||
state is None
|
||||
or state.collected
|
||||
or state.swept
|
||||
or claim is None
|
||||
or claim.operation_id != operation_id
|
||||
or claim.phase != ROUTSTR_PHASE_CLAIMED
|
||||
):
|
||||
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(CashuTransaction)
|
||||
.where(
|
||||
col(CashuTransaction.id) == state_id,
|
||||
col(CashuTransaction.request_id) == state.request_id,
|
||||
col(CashuTransaction.collected) == False, # noqa: E712
|
||||
col(CashuTransaction.swept) == False, # noqa: E712
|
||||
)
|
||||
.values(
|
||||
request_id=_routstr_request_id(
|
||||
operation_id,
|
||||
int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS,
|
||||
ROUTSTR_PHASE_SENT,
|
||||
expected_sats,
|
||||
claim.failures,
|
||||
),
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=mint_url,
|
||||
)
|
||||
)
|
||||
if (getattr(result, "rowcount", 0) or 0) != 1:
|
||||
await session.rollback()
|
||||
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
id=uuid.uuid4().hex,
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=mint_url,
|
||||
type="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _current_routstr_claim(row: UpstreamProviderRow) -> RoutstrClaim | None:
|
||||
@@ -1081,7 +1114,11 @@ async def _set_ppq_state_terminal(
|
||||
col(CashuTransaction.collected) == False, # noqa: E712
|
||||
col(CashuTransaction.swept) == False, # noqa: E712
|
||||
)
|
||||
.values(collected=collected, swept=swept)
|
||||
.values(
|
||||
collected=collected,
|
||||
swept=swept,
|
||||
created_at=int(time.time()) if collected else CashuTransaction.created_at,
|
||||
)
|
||||
)
|
||||
updated = (getattr(result, "rowcount", 0) or 0) == 1
|
||||
if updated:
|
||||
@@ -1106,8 +1143,10 @@ async def _reconcile_ppq_state(
|
||||
"""
|
||||
async with create_session() as session:
|
||||
transaction = await session.get(CashuTransaction, _ppq_state_id(row))
|
||||
if transaction is None or transaction.collected or transaction.swept:
|
||||
if transaction is None or transaction.swept:
|
||||
return False
|
||||
if transaction.collected:
|
||||
return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS
|
||||
|
||||
claim = _parse_ppq_request_id(transaction.request_id)
|
||||
if claim is None:
|
||||
@@ -1173,7 +1212,7 @@ async def _reconcile_ppq_state(
|
||||
|
||||
|
||||
async def _ppq_provider_is_claimable(
|
||||
session: AsyncSession, provider_id: int | None
|
||||
session: AsyncSession, row: UpstreamProviderRow
|
||||
) -> bool:
|
||||
"""Re-read the provider inside the claim transaction.
|
||||
|
||||
@@ -1184,10 +1223,16 @@ async def _ppq_provider_is_claimable(
|
||||
this the worker could create a claim for a provider that no longer
|
||||
exists, orphaning it forever.
|
||||
"""
|
||||
if provider_id is None:
|
||||
if row.id is None:
|
||||
return False
|
||||
current = await session.get(UpstreamProviderRow, provider_id)
|
||||
return current is not None and current.provider_type == "ppqai"
|
||||
current = await session.get(UpstreamProviderRow, row.id)
|
||||
return bool(
|
||||
current is not None
|
||||
and current.enabled
|
||||
and current.provider_type == "ppqai"
|
||||
and current.api_key == row.api_key
|
||||
and current.provider_settings == row.provider_settings
|
||||
)
|
||||
|
||||
|
||||
async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||
@@ -1198,10 +1243,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||
request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending")
|
||||
|
||||
async with create_session() as session:
|
||||
if not await _ppq_provider_is_claimable(session, row.id):
|
||||
if not await _ppq_provider_is_claimable(session, row):
|
||||
return None
|
||||
existing = await session.get(CashuTransaction, state_id)
|
||||
if existing is not None:
|
||||
if (
|
||||
existing.collected
|
||||
and not existing.swept
|
||||
and int(time.time()) - existing.created_at
|
||||
< PPQ_SETTLED_COOLDOWN_SECONDS
|
||||
):
|
||||
return None
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(CashuTransaction)
|
||||
.where(
|
||||
@@ -1232,7 +1284,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||
async with create_session() as session:
|
||||
# Same fencing as the update path: the provider must still exist
|
||||
# inside the transaction that creates the claim.
|
||||
if not await _ppq_provider_is_claimable(session, row.id):
|
||||
if not await _ppq_provider_is_claimable(session, row):
|
||||
return None
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
@@ -1453,6 +1505,26 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None
|
||||
if balance >= threshold_usd:
|
||||
return
|
||||
|
||||
# Require two low-balance reads before creating an invoice.
|
||||
confirmed_balance = await provider.get_balance()
|
||||
if (
|
||||
confirmed_balance is None
|
||||
or not math.isfinite(confirmed_balance)
|
||||
or confirmed_balance < 0
|
||||
or confirmed_balance >= threshold_usd
|
||||
):
|
||||
logger.info(
|
||||
"PPQ auto top-up aborted by balance confirmation",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"first_balance_usd": balance,
|
||||
"confirmed_balance_usd": confirmed_balance,
|
||||
"threshold_usd": threshold_usd,
|
||||
},
|
||||
)
|
||||
return
|
||||
balance = confirmed_balance
|
||||
|
||||
# Perform local pricing and owner-funds checks before asking PPQ to create
|
||||
# an invoice. The exact mint quote still has to be checked afterward, but
|
||||
# predictable local failures should not leave abandoned PPQ invoices.
|
||||
|
||||
+135
-96
@@ -1,12 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import math
|
||||
import traceback
|
||||
import typing
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Iterator
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterator
|
||||
from typing import Any, Mapping, Self, cast
|
||||
|
||||
import httpx
|
||||
@@ -70,6 +71,32 @@ if typing.TYPE_CHECKING:
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
async def _aclose_if_needed(resource: object | None) -> None:
|
||||
if resource is None:
|
||||
return
|
||||
close = getattr(resource, "aclose", None)
|
||||
if close is None:
|
||||
return
|
||||
result = close()
|
||||
if inspect.isawaitable(result):
|
||||
await result
|
||||
|
||||
|
||||
async def _finalize_and_close_stream(
|
||||
finalize: Callable[[], Awaitable[None]] | None,
|
||||
response: object | None,
|
||||
client: httpx.AsyncClient | None,
|
||||
) -> None:
|
||||
try:
|
||||
if finalize is not None:
|
||||
await finalize()
|
||||
finally:
|
||||
try:
|
||||
await _aclose_if_needed(response)
|
||||
finally:
|
||||
await _aclose_if_needed(client)
|
||||
|
||||
|
||||
CostMetadata = CostData | MaxCostData | dict[str, Any]
|
||||
|
||||
|
||||
@@ -993,6 +1020,7 @@ class BaseUpstreamProvider:
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1034,23 +1062,37 @@ class BaseUpstreamProvider:
|
||||
nonlocal usage_finalized
|
||||
if usage_finalized:
|
||||
return
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Fallback stream billing finalization failed; releasing reservation",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
usage_finalized = (
|
||||
await self._release_failed_streaming_reservation(
|
||||
fresh_key, new_session, reservation_snapshot
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Fallback stream billing recovery could not access the database",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
|
||||
def _process_event(
|
||||
raw_event: bytes, final: bool = False
|
||||
@@ -1280,18 +1322,24 @@ class BaseUpstreamProvider:
|
||||
|
||||
except Exception as stream_error:
|
||||
logger.warning(
|
||||
"Streaming interrupted; finalizing in background",
|
||||
"Streaming interrupted; finalizing before closing upstream",
|
||||
extra={
|
||||
"error": str(stream_error),
|
||||
"error_type": type(stream_error).__name__,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
if not usage_finalized:
|
||||
# Create a background task to ensure finalization happens
|
||||
# even if the generator is closed early
|
||||
background_tasks.add_task(finalize_db_only)
|
||||
# Shielded so a client disconnect cannot cancel billing
|
||||
# finalization or leak the upstream connection.
|
||||
await asyncio.shield(
|
||||
_finalize_and_close_stream(
|
||||
None if usage_finalized else finalize_db_only,
|
||||
response,
|
||||
client,
|
||||
)
|
||||
)
|
||||
|
||||
# Remove inaccurate encoding headers from upstream response
|
||||
response_headers = dict(response.headers)
|
||||
@@ -1451,6 +1499,7 @@ class BaseUpstreamProvider:
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1484,23 +1533,37 @@ class BaseUpstreamProvider:
|
||||
nonlocal usage_finalized
|
||||
if usage_finalized:
|
||||
return
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Fallback Responses billing finalization failed; releasing reservation",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
usage_finalized = (
|
||||
await self._release_failed_streaming_reservation(
|
||||
fresh_key, new_session, reservation_snapshot
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Fallback Responses billing recovery could not access the database",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
|
||||
def _process_event(
|
||||
raw_event: bytes, final: bool = False
|
||||
@@ -1690,16 +1753,24 @@ class BaseUpstreamProvider:
|
||||
|
||||
except Exception as stream_error:
|
||||
logger.warning(
|
||||
"Responses API streaming interrupted; finalizing in background",
|
||||
"Responses API streaming interrupted; finalizing before closing upstream",
|
||||
extra={
|
||||
"error": str(stream_error),
|
||||
"error_type": type(stream_error).__name__,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
if not usage_finalized:
|
||||
await finalize_db_only()
|
||||
# Shielded so a client disconnect cannot cancel billing
|
||||
# finalization or leak the upstream connection.
|
||||
await asyncio.shield(
|
||||
_finalize_and_close_stream(
|
||||
None if usage_finalized else finalize_db_only,
|
||||
response,
|
||||
client,
|
||||
)
|
||||
)
|
||||
|
||||
# Remove inaccurate encoding headers from upstream response
|
||||
response_headers = dict(response.headers)
|
||||
@@ -3051,9 +3122,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
if is_streaming and response.status_code == 200:
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
result = await self.handle_streaming_chat_completion(
|
||||
return await self.handle_streaming_chat_completion(
|
||||
response,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
@@ -3061,9 +3130,8 @@ class BaseUpstreamProvider:
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
reservation_snapshot=reservation_snapshot,
|
||||
client=client,
|
||||
)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
|
||||
# Handle both non-streaming chat completions and embeddings
|
||||
if response.status_code == 200:
|
||||
@@ -3332,19 +3400,15 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
|
||||
if is_streaming and response.status_code == 200:
|
||||
result = await self.handle_streaming_responses_completion(
|
||||
return await self.handle_streaming_responses_completion(
|
||||
response,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
reservation_snapshot=reservation_snapshot,
|
||||
client=client,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
|
||||
if response.status_code == 200:
|
||||
try:
|
||||
@@ -3628,54 +3692,30 @@ class BaseUpstreamProvider:
|
||||
extra={"amount": amount, "unit": unit, "mint": mint},
|
||||
)
|
||||
|
||||
max_retries = 3
|
||||
last_exception = None
|
||||
refund_token = None
|
||||
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
refund_token = await send_token(amount, unit=unit, mint_url=mint)
|
||||
break
|
||||
except Exception as e:
|
||||
last_exception = e
|
||||
if attempt < max_retries - 1:
|
||||
logger.warning(
|
||||
"Refund token creation failed, retrying",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"attempt": attempt + 1,
|
||||
"max_retries": max_retries,
|
||||
"amount": amount,
|
||||
"unit": unit,
|
||||
"mint": mint,
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.error(
|
||||
"Failed to create refund token after all retries",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"attempt": attempt + 1,
|
||||
"max_retries": max_retries,
|
||||
"amount": amount,
|
||||
"unit": unit,
|
||||
"mint": mint,
|
||||
},
|
||||
)
|
||||
|
||||
if refund_token is None:
|
||||
try:
|
||||
# Token creation may swap proofs, so it is unsafe to retry.
|
||||
refund_token = await send_token(amount, unit=unit, mint_url=mint)
|
||||
except Exception as error:
|
||||
logger.error(
|
||||
"Failed to create refund token",
|
||||
extra={
|
||||
"error": str(error),
|
||||
"error_type": type(error).__name__,
|
||||
"amount": amount,
|
||||
"unit": unit,
|
||||
"mint": mint,
|
||||
},
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"failed to create refund after {max_retries} attempts: {str(last_exception)}",
|
||||
"message": f"failed to create refund: {error}",
|
||||
"type": "invalid_request_error",
|
||||
"code": "send_token_failed",
|
||||
}
|
||||
},
|
||||
)
|
||||
) from error
|
||||
|
||||
logger.info(
|
||||
"Refund token created successfully",
|
||||
@@ -3683,7 +3723,6 @@ class BaseUpstreamProvider:
|
||||
"amount": amount,
|
||||
"unit": unit,
|
||||
"mint": mint,
|
||||
"attempt": attempt + 1,
|
||||
"token_preview": refund_token[:20] + "..."
|
||||
if len(refund_token) > 20
|
||||
else refund_token,
|
||||
@@ -5362,7 +5401,7 @@ class BaseUpstreamProvider:
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to refresh models cache for {self.provider_type or self.base_url}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
extra={"error": repr(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
def get_cached_models(self) -> list[Model]:
|
||||
|
||||
+180
-108
@@ -1,5 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import httpx
|
||||
@@ -15,6 +19,87 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
_PPQ_SAFE_READ_ATTEMPTS = 3
|
||||
_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0
|
||||
|
||||
|
||||
class PPQCircuitOpenError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PPQCircuitState:
|
||||
consecutive_failures: int = 0
|
||||
cooldown_until: float = 0.0
|
||||
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
loop: asyncio.AbstractEventLoop | None = None
|
||||
|
||||
|
||||
_ppq_circuits: dict[str, _PPQCircuitState] = {}
|
||||
|
||||
|
||||
def _ppq_origin(url: str) -> str:
|
||||
parsed = httpx.URL(url)
|
||||
port = parsed.port or {"https": 443, "http": 80}.get(parsed.scheme, 0)
|
||||
return f"{parsed.scheme}://{parsed.host}:{port}"
|
||||
|
||||
|
||||
async def _safe_read_request(
|
||||
client: httpx.AsyncClient,
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
headers: dict[str, str],
|
||||
json: dict[str, object] | None = None,
|
||||
) -> httpx.Response:
|
||||
state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState())
|
||||
loop = asyncio.get_running_loop()
|
||||
if state.loop is not loop:
|
||||
# Locks cannot be reused across event loops.
|
||||
state.lock = asyncio.Lock()
|
||||
state.loop = loop
|
||||
async with state.lock:
|
||||
remaining = state.cooldown_until - time.monotonic()
|
||||
if remaining > 0:
|
||||
raise PPQCircuitOpenError(
|
||||
f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s"
|
||||
)
|
||||
|
||||
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
|
||||
try:
|
||||
response = await client.request(method, url, headers=headers, json=json)
|
||||
response.raise_for_status()
|
||||
state.consecutive_failures = 0
|
||||
state.cooldown_until = 0.0
|
||||
return response
|
||||
except (httpx.TransportError, httpx.HTTPStatusError) as error:
|
||||
retryable_status = isinstance(error, httpx.HTTPStatusError) and (
|
||||
error.response.status_code in {502, 503, 504}
|
||||
)
|
||||
if not isinstance(error, httpx.TransportError) and not retryable_status:
|
||||
raise
|
||||
state.consecutive_failures += 1
|
||||
if attempt >= _PPQ_SAFE_READ_ATTEMPTS:
|
||||
state.cooldown_until = (
|
||||
time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS
|
||||
)
|
||||
raise
|
||||
base_delay = 0.25 * (2 ** (attempt - 1))
|
||||
delay = base_delay + random.uniform(0.0, base_delay)
|
||||
logger.warning(
|
||||
"PPQ.AI safe read failed; retrying",
|
||||
extra={
|
||||
"url": url,
|
||||
"attempt": attempt,
|
||||
"max_attempts": _PPQ_SAFE_READ_ATTEMPTS,
|
||||
"backoff_seconds": round(delay, 3),
|
||||
"error": repr(error),
|
||||
"error_type": type(error).__name__,
|
||||
},
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
raise RuntimeError("unreachable")
|
||||
|
||||
|
||||
class PPQAIModelPricing(BaseModel):
|
||||
ui: Optional[dict[str, float]] = None
|
||||
@@ -123,122 +208,109 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
url = f"{self.base_url}/models"
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await _safe_read_request(client, "GET", url, headers=headers)
|
||||
data = response.json()
|
||||
|
||||
models_data = data.get("data", [])
|
||||
models_data = data.get("data", [])
|
||||
|
||||
or_models = [
|
||||
Model(**model) # type: ignore
|
||||
for model in await async_fetch_openrouter_models()
|
||||
]
|
||||
or_models = [
|
||||
Model(**model) # type: ignore
|
||||
for model in await async_fetch_openrouter_models()
|
||||
]
|
||||
|
||||
models = []
|
||||
for model_data in models_data:
|
||||
try:
|
||||
ppqai_model = PPQAIModel.parse_obj(model_data)
|
||||
if ppqai_model.id in self.IGNORED_MODEL_IDS:
|
||||
continue
|
||||
models = []
|
||||
for model_data in models_data:
|
||||
try:
|
||||
ppqai_model = PPQAIModel.parse_obj(model_data)
|
||||
if ppqai_model.id in self.IGNORED_MODEL_IDS:
|
||||
continue
|
||||
|
||||
or_model = next(
|
||||
(
|
||||
model
|
||||
for model in or_models
|
||||
if (model.id == ppqai_model.id)
|
||||
or (model.id.split("/")[-1] == ppqai_model.id)
|
||||
or (model.id == ppqai_model.id.split("/")[-1])
|
||||
),
|
||||
None,
|
||||
)
|
||||
or_model = next(
|
||||
(
|
||||
model
|
||||
for model in or_models
|
||||
if (model.id == ppqai_model.id)
|
||||
or (model.id.split("/")[-1] == ppqai_model.id)
|
||||
or (model.id == ppqai_model.id.split("/")[-1])
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
if or_model:
|
||||
input_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M"
|
||||
)
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
if or_model:
|
||||
input_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get("input_per_1M")
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
|
||||
if input_price is not None:
|
||||
or_model.pricing.prompt = input_price / 1_000_000
|
||||
if input_price is not None:
|
||||
or_model.pricing.prompt = input_price / 1_000_000
|
||||
|
||||
output_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M"
|
||||
)
|
||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||
output_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get("output_per_1M")
|
||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||
|
||||
if output_price is not None:
|
||||
or_model.pricing.completion = output_price / 1_000_000
|
||||
if output_price is not None:
|
||||
or_model.pricing.completion = output_price / 1_000_000
|
||||
|
||||
if cl := ppqai_model.context_length:
|
||||
or_model.context_length = cl
|
||||
models.append(or_model)
|
||||
else:
|
||||
input_price = 0.0
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M", 0.0
|
||||
)
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
|
||||
output_price = 0.0
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M", 0.0
|
||||
)
|
||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||
|
||||
models.append(
|
||||
Model(
|
||||
id=ppqai_model.id,
|
||||
name=ppqai_model.name,
|
||||
created=ppqai_model.created_at // 1000,
|
||||
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
|
||||
context_length=ppqai_model.context_length,
|
||||
architecture=Architecture(
|
||||
modality="text->text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="Unknown",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(
|
||||
prompt=input_price / 1_000_000,
|
||||
completion=output_price / 1_000_000,
|
||||
request=0.0,
|
||||
image=0.0,
|
||||
web_search=0.0,
|
||||
internal_reasoning=0.0,
|
||||
),
|
||||
)
|
||||
if cl := ppqai_model.context_length:
|
||||
or_model.context_length = cl
|
||||
models.append(or_model)
|
||||
else:
|
||||
input_price = 0.0
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M", 0.0
|
||||
)
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
|
||||
output_price = 0.0
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M", 0.0
|
||||
)
|
||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||
|
||||
models.append(
|
||||
Model(
|
||||
id=ppqai_model.id,
|
||||
name=ppqai_model.name,
|
||||
created=ppqai_model.created_at // 1000,
|
||||
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
|
||||
context_length=ppqai_model.context_length,
|
||||
architecture=Architecture(
|
||||
modality="text->text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="Unknown",
|
||||
instruct_type=None,
|
||||
),
|
||||
pricing=Pricing(
|
||||
prompt=input_price / 1_000_000,
|
||||
completion=output_price / 1_000_000,
|
||||
request=0.0,
|
||||
image=0.0,
|
||||
web_search=0.0,
|
||||
internal_reasoning=0.0,
|
||||
),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to parse PPQ.AI model",
|
||||
extra={
|
||||
"model_id": model_data.get("id", "unknown"),
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to parse PPQ.AI model",
|
||||
extra={
|
||||
"model_id": model_data.get("id", "unknown"),
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
|
||||
return models
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error fetching models from PPQ.AI",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
return []
|
||||
return models
|
||||
|
||||
async def on_upstream_error_redirect(
|
||||
self, status_code: int, error_message: str
|
||||
@@ -360,8 +432,7 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
response = await _safe_read_request(client, "GET", url, headers=headers)
|
||||
status_data = response.json()
|
||||
|
||||
is_paid = status_data.get("status") == "Settled"
|
||||
@@ -460,8 +531,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
logger.debug("Checking PPQ.AI account balance", extra={"url": url})
|
||||
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.post(url, headers=headers, json={})
|
||||
response.raise_for_status()
|
||||
response = await _safe_read_request(
|
||||
client, "POST", url, headers=headers, json={}
|
||||
)
|
||||
balance_data = response.json()
|
||||
|
||||
logger.debug(
|
||||
|
||||
+109
-30
@@ -14,6 +14,7 @@ from typing import AsyncGenerator, TypedDict
|
||||
import httpx
|
||||
from cashu.core.base import MeltQuote, MeltQuoteState, MintQuote, Proof, Token
|
||||
from cashu.core.mint_info import MintInfo as _CashuMintInfo
|
||||
from cashu.wallet.crud import get_keysets as get_cashu_keysets
|
||||
from cashu.wallet.helpers import deserialize_token_from_string
|
||||
from cashu.wallet.wallet import Wallet as _CashuWallet
|
||||
from pydantic_core import PydanticUndefined
|
||||
@@ -121,6 +122,12 @@ def _mints_to_inspect() -> list[str]:
|
||||
return mint_urls
|
||||
|
||||
|
||||
_WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS = 30
|
||||
_WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS = 300
|
||||
_mint_metadata_last_load: dict[str, float] = {}
|
||||
_mint_metadata_load_locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
|
||||
class Wallet(_CashuWallet):
|
||||
"""Cashu adapter that preserves HTTP 429 for Routstr's mint policy."""
|
||||
|
||||
@@ -141,11 +148,34 @@ class Wallet(_CashuWallet):
|
||||
_CashuWallet.raise_on_error_request(resp)
|
||||
|
||||
async def load_mint(
|
||||
self, keyset_id: str = "", force_old_keysets: bool = False
|
||||
self,
|
||||
keyset_id: str = "",
|
||||
force_old_keysets: bool = False,
|
||||
*,
|
||||
force_refresh: bool = False,
|
||||
) -> None:
|
||||
await self.load_mint_keysets(force_old_keysets)
|
||||
await self.activate_keyset(keyset_id)
|
||||
await self.load_mint_info(reload=True)
|
||||
mint_url = str(self.url)
|
||||
lock = _mint_metadata_load_locks.setdefault(mint_url, asyncio.Lock())
|
||||
async with lock:
|
||||
now = time.monotonic()
|
||||
last = _mint_metadata_last_load.get(mint_url)
|
||||
if (
|
||||
not force_refresh
|
||||
and last is not None
|
||||
and now - last < _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
|
||||
):
|
||||
try:
|
||||
await self.load_keysets_from_db()
|
||||
await self.activate_keyset(keyset_id)
|
||||
await self.load_mint_info(reload=False)
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await self.load_mint_keysets(force_old_keysets)
|
||||
await self.activate_keyset(keyset_id)
|
||||
await self.load_mint_info(reload=True)
|
||||
_mint_metadata_last_load[mint_url] = time.monotonic()
|
||||
|
||||
|
||||
class MintConnectionError(Exception):
|
||||
@@ -491,7 +521,11 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
|
||||
|
||||
|
||||
async def _send_locked(
|
||||
amount: int, unit: str, mint_url: str | None = None
|
||||
amount: int,
|
||||
unit: str,
|
||||
mint_url: str | None = None,
|
||||
*,
|
||||
owner_only: bool = False,
|
||||
) -> tuple[int, str]:
|
||||
effective_mint_url = await find_trusted_mint_with_funds(
|
||||
amount, unit, mint_url, force_reload=True
|
||||
@@ -501,6 +535,12 @@ async def _send_locked(
|
||||
wallet, effective_mint_url, unit, not_reserved=True
|
||||
)
|
||||
proofs_for_mint = sum(proof.amount for proof in proofs)
|
||||
if owner_only:
|
||||
owner_balance = await _owner_balance_for_mint_and_unit(
|
||||
effective_mint_url, unit, proofs_for_mint
|
||||
)
|
||||
if owner_balance < amount:
|
||||
raise ValueError("Owner Cashu balance is insufficient for auto top-up")
|
||||
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
|
||||
|
||||
@@ -547,6 +587,13 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str
|
||||
return token
|
||||
|
||||
|
||||
async def send_token_from_owner_locked(
|
||||
amount: int, unit: str, mint_url: str | None = None
|
||||
) -> str:
|
||||
_, token = await _send_locked(amount, unit, mint_url, owner_only=True)
|
||||
return token
|
||||
|
||||
|
||||
class Bolt11PaymentNotAttempted(Exception):
|
||||
"""The invoice was definitively not paid, so the attempt can be retried.
|
||||
|
||||
@@ -1048,6 +1095,7 @@ async def _request_mint_with_fallback(
|
||||
mint_url,
|
||||
settings.primary_mint_unit,
|
||||
retry_on_rate_limit=False,
|
||||
load_proofs=False,
|
||||
)
|
||||
quote = await run_mint_operation(
|
||||
lambda: wallet.request_mint(amount),
|
||||
@@ -1179,6 +1227,7 @@ async def _calculate_swap_amount(
|
||||
lambda: token_wallet.melt_quote(dummy_mint_quote.request),
|
||||
op_name="swap_fee_est_melt_quote",
|
||||
mint_url=token_mint_url,
|
||||
retry_timeouts=False,
|
||||
)
|
||||
|
||||
fee_reserve = dummy_melt_quote.fee_reserve
|
||||
@@ -1416,6 +1465,7 @@ async def swap_to_trusted_mint(
|
||||
lambda: token_wallet.melt_quote(mint_quote.request),
|
||||
op_name="swap_melt_quote",
|
||||
mint_url=token_obj.mint,
|
||||
retry_timeouts=False,
|
||||
)
|
||||
except Exception as error:
|
||||
if is_mint_connection_error(error):
|
||||
@@ -1788,21 +1838,34 @@ async def _credit_balance_locked(
|
||||
)
|
||||
return amount
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"credit_balance: Error during token redemption",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
classification = classify_redemption_error(e)
|
||||
expected_codes = {
|
||||
"cashu_token_already_spent",
|
||||
"cashu_source_mint_unreachable",
|
||||
"cashu_mint_unreachable",
|
||||
"cashu_mint_rate_limited",
|
||||
}
|
||||
log = (
|
||||
logger.info
|
||||
if classification is not None and classification[3] in expected_codes
|
||||
else logger.error
|
||||
)
|
||||
log(
|
||||
"credit_balance: Token redemption failed",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"error_code": classification[3] if classification else None,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
_wallets: dict[str, Wallet] = {}
|
||||
# Proofs require a shorter refresh interval than remote mint metadata.
|
||||
_wallet_last_load: dict[str, float] = {}
|
||||
_wallet_last_mint_load: dict[str, float] = {}
|
||||
_wallet_load_locks: dict[str, asyncio.Lock] = {}
|
||||
# Minimum seconds between full mint info + proof reloads for the same
|
||||
# wallet. Prevents redundant mint API calls when get_wallet(load=True)
|
||||
# is called rapidly by multiple background tasks (balance fetch, payout,
|
||||
# auto-topup all hitting get_wallet within the same cycle).
|
||||
_WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30
|
||||
|
||||
|
||||
async def get_wallet(
|
||||
@@ -1811,8 +1874,9 @@ async def get_wallet(
|
||||
load: bool = True,
|
||||
retry_on_rate_limit: bool = True,
|
||||
force_reload: bool = False,
|
||||
load_proofs: bool = True,
|
||||
) -> Wallet:
|
||||
global _wallets, _wallet_last_load, _wallet_load_locks
|
||||
global _wallets, _wallet_last_load, _wallet_last_mint_load, _wallet_load_locks
|
||||
id = f"{mint_url}_{unit}"
|
||||
lock = _wallet_load_locks.setdefault(id, asyncio.Lock())
|
||||
async with lock:
|
||||
@@ -1821,25 +1885,39 @@ async def get_wallet(
|
||||
|
||||
if load:
|
||||
now = time.monotonic()
|
||||
last = _wallet_last_load.get(id)
|
||||
last_mint_load = _wallet_last_mint_load.get(id)
|
||||
if (
|
||||
force_reload
|
||||
or last is None
|
||||
or now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS
|
||||
or last_mint_load is None
|
||||
or now - last_mint_load >= _WALLET_MINT_RELOAD_MIN_INTERVAL_SECONDS
|
||||
):
|
||||
await run_mint_operation(
|
||||
lambda: _wallets[id].load_mint(),
|
||||
lambda: (
|
||||
_wallets[id].load_mint(force_refresh=True)
|
||||
if force_reload
|
||||
else _wallets[id].load_mint()
|
||||
),
|
||||
op_name="load_mint",
|
||||
mint_url=mint_url,
|
||||
retry_on_rate_limit=retry_on_rate_limit,
|
||||
)
|
||||
await run_mint_operation(
|
||||
lambda: _wallets[id].load_proofs(reload=True),
|
||||
op_name="load_proofs",
|
||||
mint_url=mint_url,
|
||||
retry_on_rate_limit=retry_on_rate_limit,
|
||||
)
|
||||
_wallet_last_load[id] = time.monotonic()
|
||||
_wallet_last_mint_load[id] = time.monotonic()
|
||||
|
||||
if load_proofs:
|
||||
last_proof_load = _wallet_last_load.get(id)
|
||||
if (
|
||||
force_reload
|
||||
or last_proof_load is None
|
||||
or now - last_proof_load
|
||||
>= _WALLET_PROOF_RELOAD_MIN_INTERVAL_SECONDS
|
||||
):
|
||||
await run_mint_operation(
|
||||
lambda: _wallets[id].load_proofs(reload=True),
|
||||
op_name="load_proofs",
|
||||
mint_url=mint_url,
|
||||
retry_on_rate_limit=retry_on_rate_limit,
|
||||
)
|
||||
_wallet_last_load[id] = time.monotonic()
|
||||
return _wallets[id]
|
||||
|
||||
|
||||
@@ -1912,13 +1990,14 @@ async def _get_supported_mint_units(mint_url: str) -> list[str]:
|
||||
if cached is not None and now < cached[0]:
|
||||
return cached[1]
|
||||
|
||||
wallet = await get_wallet(mint_url, settings.primary_mint_unit, load=False)
|
||||
keysets = await run_mint_operation(
|
||||
lambda: wallet._get_keysets(),
|
||||
op_name="get_mint_keysets",
|
||||
mint_url=mint_url,
|
||||
# A metadata load populates Cashu's shared keyset cache for all units.
|
||||
wallet = await get_wallet(
|
||||
mint_url,
|
||||
settings.primary_mint_unit,
|
||||
retry_on_rate_limit=False,
|
||||
load_proofs=False,
|
||||
)
|
||||
keysets = await get_cashu_keysets(mint_url=wallet.url, db=wallet.db)
|
||||
units: list[str] = []
|
||||
for keyset in keysets:
|
||||
if not keyset.active or keyset.unit is None:
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
@@ -34,11 +35,14 @@ async def test_structured_http_error_uses_standard_error_envelope() -> None:
|
||||
"details": {"mint": "https://mint.example"},
|
||||
}
|
||||
|
||||
response = await http_exception_handler(
|
||||
request,
|
||||
HTTPException(status_code=503, detail={"error": error}),
|
||||
)
|
||||
with patch("routstr.core.exceptions.logger") as logger:
|
||||
response = await http_exception_handler(
|
||||
request,
|
||||
HTTPException(status_code=503, detail={"error": error}),
|
||||
)
|
||||
|
||||
logger.warning.assert_called_once()
|
||||
logger.error.assert_not_called()
|
||||
assert response.status_code == 503
|
||||
assert json.loads(response.body) == {
|
||||
"detail": {"error": error},
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.melt_quote = AsyncMock(
|
||||
return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount)
|
||||
)
|
||||
wallet.get_fees_for_proofs.return_value = 0
|
||||
if quote_amount is None:
|
||||
wallet.melt_quote = AsyncMock(
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=1, quote="q", amount=1000),
|
||||
MagicMock(fee_reserve=1, quote="q", amount=EXPECTED_QUOTE_SAT),
|
||||
]
|
||||
)
|
||||
else:
|
||||
wallet.melt_quote = AsyncMock(
|
||||
return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount)
|
||||
)
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
wallet.get_fees_for_proofs = MagicMock(return_value=0)
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
return wallet, proofs
|
||||
@@ -124,16 +134,21 @@ async def test_raw_send_to_lnurl_accepts_exact_invoice() -> None:
|
||||
)
|
||||
|
||||
assert paid == EXPECTED_QUOTE_SAT * 1000
|
||||
wallet.select_to_send.assert_awaited_once()
|
||||
wallet.select_to_send.assert_not_awaited()
|
||||
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=True)
|
||||
wallet.melt.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None:
|
||||
# 1_000_000 msat minus an 11 sat fee reserve leaves 989_000 msat.
|
||||
wallet, proofs = _wallet(989_000)
|
||||
wallet, proofs = _wallet()
|
||||
proofs[0].amount = 1_000_000
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
wallet.melt_quote = AsyncMock(
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=1, quote="q", amount=1_000_000),
|
||||
MagicMock(fee_reserve=1, quote="q", amount=999_999),
|
||||
]
|
||||
)
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with data_patch, invoice_patch:
|
||||
@@ -141,7 +156,46 @@ async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None:
|
||||
wallet, proofs, "owner@ln.tld", "msat", amount=1_000_000
|
||||
)
|
||||
|
||||
assert paid == 989_000
|
||||
assert paid == 999_999
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion() -> (
|
||||
None
|
||||
):
|
||||
proofs = [MagicMock(amount=1, reserved=False) for _ in range(1500)]
|
||||
wallet = MagicMock(url="https://mint.test")
|
||||
wallet.get_fees_for_proofs = MagicMock(
|
||||
side_effect=lambda selected: math.ceil(len(selected) / 100)
|
||||
)
|
||||
wallet.melt_quote = AsyncMock(
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=10, quote="q1", amount=1500),
|
||||
MagicMock(fee_reserve=10, quote="q2", amount=1475),
|
||||
]
|
||||
)
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
checkpoint = AsyncMock()
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with data_patch, invoice_patch:
|
||||
paid = await raw_send_to_lnurl(
|
||||
wallet,
|
||||
proofs,
|
||||
"owner@ln.tld",
|
||||
"sat",
|
||||
amount=1500,
|
||||
on_melt_quote=checkpoint,
|
||||
)
|
||||
|
||||
assert paid == 1_475_000
|
||||
assert wallet.melt_quote.await_count == 2
|
||||
checkpoint.assert_awaited_once_with("q2")
|
||||
wallet.select_to_send.assert_not_called()
|
||||
selected = wallet.melt.await_args.kwargs["proofs"]
|
||||
assert sum(proof.amount for proof in selected) == 1500
|
||||
assert 1475 + 10 + wallet.get_fees_for_proofs(selected) == 1500
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -291,3 +345,23 @@ async def test_send_to_lnurl_does_not_reserve_before_lnurl_validation() -> None:
|
||||
assert raw_send.await_args is not None
|
||||
assert raw_send.await_args.args[1] is proofs
|
||||
assert raw_send.await_args.kwargs["amount"] == 1000
|
||||
|
||||
|
||||
def test_select_melt_proofs_stops_at_minimal_cover_when_over_budget() -> None:
|
||||
from routstr.payment.lnurl import _select_melt_proofs
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_fees_for_proofs = MagicMock(side_effect=lambda selected: len(selected))
|
||||
proofs = [MagicMock(amount=600, reserved=False) for _ in range(3)]
|
||||
|
||||
selected, shortfall = _select_melt_proofs(
|
||||
wallet,
|
||||
proofs,
|
||||
quote_amount=1000,
|
||||
fee_reserve=0,
|
||||
gross_budget=1000,
|
||||
)
|
||||
|
||||
assert selected is None
|
||||
assert shortfall == 2
|
||||
assert wallet.get_fees_for_proofs.call_count == 2
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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] = []
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user