mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-06 04:38:22 +00:00
fix: harden melt sizing, mint cooldowns, PPQ reads, and streaming finalization
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
|
||||
# Seconds a file-backed SQLite writer waits for the write lock (default: 30).
|
||||
# DATABASE_BUSY_TIMEOUT=30
|
||||
# SQLite serialises writes; increasing its pool can trade pool timeouts for
|
||||
# "database is locked" errors rather than increasing write throughput.
|
||||
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
# ADR-001: Cashu payment safety boundaries
|
||||
|
||||
## Status
|
||||
|
||||
Accepted
|
||||
|
||||
## Context
|
||||
|
||||
Cashu proofs are bearer instruments. Retrying quote creation, refund delivery, or a dispatched Lightning melt can duplicate side effects or spend proofs whose outcome is still unknown. Mint transport failures and concurrent workers also need one shared policy.
|
||||
|
||||
## Decision
|
||||
|
||||
- Treat account, invoice, quote, melt, token-delivery, and refund creation as non-idempotent unless an upstream idempotency key is available.
|
||||
- A dispatched melt with an unknown outcome keeps a durable quote-linked proof reservation until later reconciliation confirms a terminal state. An immediate `unpaid` observation after transport loss is not terminal.
|
||||
- Size melts from the quote amount, reserve, and exact proof input fees within the caller's gross budget; do not use recursive send selection for melt planning.
|
||||
- Apply mint transport/rate cooldowns centrally and permit only explicit reconciliation probes during cooldown.
|
||||
- Auto-topups require fresh threshold confirmation, durable per-provider claims/cooldown, atomic spend-cap checks, and owner-only funds.
|
||||
|
||||
## Consequences
|
||||
|
||||
Transient failures can delay payouts/topups rather than risk duplicate payment. Operators may need to reconcile ambiguous claims. Tests must cover restart, concurrency, and partial-stream failures at these boundaries.
|
||||
+26
-5
@@ -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,
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
@@ -40,15 +40,27 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
|
||||
path = request.url.path
|
||||
|
||||
# 4xx is client behaviour; the uvicorn access log already records it.
|
||||
# Only 5xx warrants a server-side warning/error log here.
|
||||
# Retryable mint outages are expected dependency failures, not application
|
||||
# faults, so keep them visible without flooding the error stream.
|
||||
if status_code >= 500:
|
||||
logger.error(
|
||||
error_type = None
|
||||
if isinstance(detail, dict):
|
||||
error = detail.get("error")
|
||||
if isinstance(error, dict):
|
||||
error_type = error.get("type")
|
||||
log = (
|
||||
logger.warning
|
||||
if error_type in {"mint_unreachable", "mint_rate_limited"}
|
||||
else logger.error
|
||||
)
|
||||
log(
|
||||
f"HTTP {status_code} on {path}: {detail}",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"status_code": status_code,
|
||||
"detail": detail,
|
||||
"path": path,
|
||||
"error_type": error_type,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
+24
-1
@@ -177,6 +177,8 @@ class MintRateGuard:
|
||||
if isinstance(error, httpx.HTTPStatusError):
|
||||
retry_after = parse_retry_after(error.response.headers)
|
||||
self.apply_rate_limit_cooldown(retry_after)
|
||||
elif is_mint_transport_error(error):
|
||||
self.apply_cooldown(MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport")
|
||||
else:
|
||||
self.apply_cooldown(1.0)
|
||||
logger.warning(
|
||||
@@ -236,6 +238,18 @@ def mint_cooldown_reason(mint_url: str) -> str | None:
|
||||
return MintRateGuard.get(mint_url).cooldown_reason()
|
||||
|
||||
|
||||
def is_mint_transport_error(error: BaseException) -> bool:
|
||||
"""Return whether an exception chain contains a mint transport failure."""
|
||||
current: BaseException | None = error
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
if isinstance(current, MINT_TRANSPORT_EXCEPTIONS):
|
||||
return True
|
||||
current = current.__cause__ or current.__context__
|
||||
return False
|
||||
|
||||
|
||||
def is_mint_rate_limited(error: BaseException) -> bool:
|
||||
"""Return whether an exception chain represents HTTP 429/cooldown."""
|
||||
|
||||
@@ -269,6 +283,7 @@ async def run_mint_operation(
|
||||
mint_url: str = "",
|
||||
retry_timeouts: bool = True,
|
||||
retry_on_rate_limit: bool = True,
|
||||
allow_during_cooldown: bool = False,
|
||||
) -> Any:
|
||||
"""Run one mint operation with bounded concurrency and adaptive cooldown."""
|
||||
|
||||
@@ -282,7 +297,7 @@ async def run_mint_operation(
|
||||
return await factory()
|
||||
|
||||
async def invoke() -> Any:
|
||||
if guard is not None:
|
||||
if guard is not None and not allow_during_cooldown:
|
||||
return await guard.run(timed_factory)
|
||||
return await timed_factory()
|
||||
|
||||
@@ -292,6 +307,10 @@ async def run_mint_operation(
|
||||
except MintCooldownError:
|
||||
raise
|
||||
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
|
||||
if guard is not None:
|
||||
guard.apply_cooldown(
|
||||
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
|
||||
)
|
||||
if retry_timeouts and attempt < max_attempts - 1:
|
||||
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
||||
logger.warning(
|
||||
@@ -310,6 +329,10 @@ async def run_mint_operation(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
if not is_mint_rate_limited(exc):
|
||||
if guard is not None and is_mint_transport_error(exc):
|
||||
guard.apply_cooldown(
|
||||
MINT_TRANSPORT_COOLDOWN_SECONDS, reason="transport"
|
||||
)
|
||||
raise
|
||||
|
||||
backoff = (2**attempt) + (time.monotonic() % 1.0)
|
||||
|
||||
+98
-31
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import math
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, TypedDict
|
||||
|
||||
@@ -238,6 +237,35 @@ async def get_lnurl_invoice(
|
||||
return invoice_data["pr"], invoice_data
|
||||
|
||||
|
||||
def _select_melt_proofs(
|
||||
wallet: Wallet,
|
||||
proofs: list[Proof],
|
||||
*,
|
||||
quote_amount: int,
|
||||
fee_reserve: int,
|
||||
gross_budget: int,
|
||||
) -> tuple[list[Proof] | None, int]:
|
||||
"""Select proofs that cover the quote and exact NUT-02 input fees.
|
||||
|
||||
Cashu 0.20's ``select_to_send`` may recursively swap when asked to spend a
|
||||
wallet's full balance. Melts accept overpayment and return change, so a
|
||||
bounded, largest-first selection is both safer and minimizes input fees.
|
||||
"""
|
||||
selected: list[Proof] = []
|
||||
selected_amount = 0
|
||||
required = quote_amount + fee_reserve
|
||||
for proof in sorted(proofs, key=lambda item: item.amount, reverse=True):
|
||||
if getattr(proof, "reserved", False) is True:
|
||||
continue
|
||||
selected.append(proof)
|
||||
selected_amount += proof.amount
|
||||
input_fees = int(wallet.get_fees_for_proofs(selected))
|
||||
required = quote_amount + fee_reserve + input_fees
|
||||
if required <= gross_budget and selected_amount >= required:
|
||||
return selected, 0
|
||||
return None, max(1, required - min(selected_amount, gross_budget))
|
||||
|
||||
|
||||
async def raw_send_to_lnurl(
|
||||
wallet: Wallet,
|
||||
proofs: list[Proof],
|
||||
@@ -293,38 +321,55 @@ async def raw_send_to_lnurl(
|
||||
f"({min_sendable_sat} - {max_sendable_sat} {unit})"
|
||||
)
|
||||
|
||||
estimated_fees_sat = int(max(math.ceil((amount_msat / 1000) * 0.01), 2)) + 1
|
||||
estimated_fees_msat = estimated_fees_sat * 1000
|
||||
final_amount = amount_msat - estimated_fees_msat
|
||||
# Start at the caller's gross budget and converge downward from the mint's
|
||||
# exact reserve plus NUT-02 input fees. Starting below the budget with a
|
||||
# percentage heuristic silently underpays even when the exact fees are tiny.
|
||||
final_amount = amount_msat
|
||||
|
||||
bolt11_invoice, _ = await get_lnurl_invoice(
|
||||
lnurl_data["callback_url"], final_amount
|
||||
)
|
||||
|
||||
melt_quote_resp = await run_mint_operation(
|
||||
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
||||
op_name="lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
# Creating another quote after response loss only abandons the first.
|
||||
retry_timeouts=False,
|
||||
)
|
||||
|
||||
# The invoice comes from the LNURL service, so its amount is untrusted. The
|
||||
# melt quote is the mint's own reading of it, and it must match what we
|
||||
# asked to send. Checked before the checkpoint and before reserving, so a
|
||||
# mismatch leaves no durable state and no locked proofs behind.
|
||||
quoted_amount = int(melt_quote_resp.amount)
|
||||
expected_amount = final_amount // 1000 if unit == "sat" else final_amount
|
||||
if quoted_amount != expected_amount:
|
||||
raise LNURLError(
|
||||
f"LNURL invoice amount does not match the requested amount "
|
||||
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
|
||||
selected_proofs: list[Proof] | None = None
|
||||
# Fee reserves can change with the invoice amount. Each quote reduces the
|
||||
# candidate by its exact shortfall, so this bounded fixed-point search keeps
|
||||
# the largest amount the gross budget can fund without Cashu coin selection.
|
||||
for _ in range(8):
|
||||
if final_amount < lnurl_data["min_sendable"]:
|
||||
raise LNURLError("Cashu melt fees leave no payable LNURL amount")
|
||||
bolt11_invoice, _ = await get_lnurl_invoice(
|
||||
lnurl_data["callback_url"], final_amount
|
||||
)
|
||||
melt_quote_resp = await run_mint_operation(
|
||||
lambda: wallet.melt_quote(invoice=bolt11_invoice),
|
||||
op_name="lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
# Creating another quote after response loss only abandons the first.
|
||||
retry_timeouts=False,
|
||||
)
|
||||
|
||||
quoted_amount = int(melt_quote_resp.amount)
|
||||
expected_amount = final_amount // 1000 if unit == "sat" else final_amount
|
||||
if quoted_amount != expected_amount:
|
||||
raise LNURLError(
|
||||
f"LNURL invoice amount does not match the requested amount "
|
||||
f"(quoted {quoted_amount} {unit}, expected {expected_amount} {unit})"
|
||||
)
|
||||
|
||||
selected_proofs, shortfall = _select_melt_proofs(
|
||||
wallet,
|
||||
proofs,
|
||||
quote_amount=quoted_amount,
|
||||
fee_reserve=int(melt_quote_resp.fee_reserve),
|
||||
gross_budget=amount,
|
||||
)
|
||||
if selected_proofs is not None:
|
||||
break
|
||||
final_amount -= shortfall * (1000 if unit == "sat" else 1)
|
||||
else:
|
||||
raise LNURLError("Cashu melt fees exceed the requested gross amount")
|
||||
|
||||
if on_melt_quote is not None:
|
||||
await on_melt_quote(melt_quote_resp.quote)
|
||||
|
||||
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
|
||||
proofs = selected_proofs
|
||||
await wallet.set_reserved_for_send(proofs, reserved=True)
|
||||
|
||||
try:
|
||||
melt_response = await run_mint_operation(
|
||||
@@ -365,8 +410,14 @@ async def raw_send_to_lnurl(
|
||||
else:
|
||||
melt_error = None
|
||||
|
||||
if getattr(melt_response, "state", None) == MeltQuoteState.paid:
|
||||
melt_state = getattr(melt_response, "state", None)
|
||||
if melt_state == MeltQuoteState.paid:
|
||||
return final_amount
|
||||
if melt_state == MeltQuoteState.unpaid:
|
||||
# A direct unpaid response is authoritative: the mint rejected the melt
|
||||
# and Cashu has already cleared its quote-linked reservation.
|
||||
await wallet.set_reserved_for_send(proofs, reserved=False)
|
||||
raise LNURLError("Cashu mint confirmed that the melt was unpaid")
|
||||
|
||||
try:
|
||||
quote = await run_mint_operation(
|
||||
@@ -374,6 +425,9 @@ async def raw_send_to_lnurl(
|
||||
op_name="reconcile_lnurl_melt_quote",
|
||||
mint_url=str(wallet.url),
|
||||
retry_timeouts=False,
|
||||
# One direct state lookup is required to reconcile the just-dispatched
|
||||
# melt even though its transport failure opened the mint cooldown.
|
||||
allow_during_cooldown=True,
|
||||
)
|
||||
except Exception as reconciliation_error:
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
@@ -384,9 +438,22 @@ async def raw_send_to_lnurl(
|
||||
if quote is not None and quote.state == MeltQuoteState.paid:
|
||||
return final_amount
|
||||
if quote is not None and quote.state == MeltQuoteState.unpaid:
|
||||
# get_melt_quote() has authoritatively released the melt reservation;
|
||||
# callers may restore their debit and retry with a new payment plan.
|
||||
raise LNURLError("Cashu mint confirmed that the melt was unpaid") from melt_error
|
||||
# Reaching reconciliation means melt was dispatched and either lost its
|
||||
# response or returned pending. A just-dispatched quote can briefly read
|
||||
# UNPAID before transitioning. Cashu clears the reservation while
|
||||
# refreshing that state, so restore it and require later reconciliation.
|
||||
try:
|
||||
await wallet.set_reserved_for_melt(
|
||||
proofs, reserved=True, quote_id=melt_quote_resp.quote
|
||||
)
|
||||
except Exception as reservation_error:
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous and its proof reservation could not "
|
||||
"be restored; proofs must not be retried"
|
||||
) from reservation_error
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
"Melt outcome is ambiguous; an immediate unpaid state is not final"
|
||||
) from melt_error
|
||||
|
||||
state = getattr(getattr(quote, "state", None), "value", "unknown")
|
||||
raise MeltOutcomeAmbiguousError(
|
||||
|
||||
+162
-81
@@ -15,9 +15,6 @@ from ..core.db import (
|
||||
UpstreamProviderRow,
|
||||
create_session,
|
||||
)
|
||||
from ..core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from ..payment.price import sats_usd_price
|
||||
from ..wallet import (
|
||||
Bolt11PaymentAmbiguous,
|
||||
@@ -27,7 +24,7 @@ from ..wallet import (
|
||||
maximum_owner_cashu_balance_sats,
|
||||
prepare_bolt11_payment,
|
||||
release_token_reservation,
|
||||
send_token,
|
||||
send_token_from_owner_locked,
|
||||
token_mint_url,
|
||||
wallet_operation_guard,
|
||||
)
|
||||
@@ -49,6 +46,7 @@ PPQ_PHASES = frozenset({PPQ_PHASE_CLAIMED, PPQ_PHASE_IN_FLIGHT, PPQ_PHASE_RECONC
|
||||
PPQ_SETTLEMENT_ATTEMPTS = 5
|
||||
PPQ_SETTLEMENT_POLL_SECONDS = 2
|
||||
PPQ_PENDING_TTL_SECONDS = 15 * 60
|
||||
PPQ_SETTLED_COOLDOWN_SECONDS = 5 * 60
|
||||
PPQ_MAX_INVOICE_PREMIUM = 1.10
|
||||
PPQ_MIN_TOPUP_USD = 1
|
||||
PPQ_MAX_TOPUP_USD = 500
|
||||
@@ -99,7 +97,7 @@ async def periodic_auto_topup() -> None:
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Auto top-up cycle failed",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
extra={"error": repr(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS)
|
||||
@@ -130,7 +128,7 @@ async def _run_auto_topup_cycle() -> None:
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"base_url": row.base_url,
|
||||
"error": str(e),
|
||||
"error": repr(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
@@ -161,7 +159,11 @@ async def _reconcile_all_ppq_claims() -> set[int]:
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"PPQ claim reconciliation failed",
|
||||
extra={"provider_id": row.id, "error": str(e)},
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"error": repr(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
return active_provider_ids
|
||||
|
||||
@@ -363,67 +365,64 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
)
|
||||
|
||||
try:
|
||||
token = await send_token(amount, "sat", mint_url)
|
||||
async with wallet_operation_guard():
|
||||
# The cap, owner-liability check, proof reservation, and outgoing
|
||||
# audit row share one wallet mutation scope. The audit row must be
|
||||
# durable before another worker can recheck the rolling cap.
|
||||
spent_24h_sats = await _routstr_spent_last_24h_sats()
|
||||
if spent_24h_sats + amount > ROUTSTR_MAX_DAILY_TOPUP_SATS:
|
||||
raise ValueError("Routstr auto top-up daily spend cap reached")
|
||||
token = await send_token_from_owner_locked(amount, "sat", mint_url)
|
||||
actual_mint_url = token_mint_url(token, mint_url)
|
||||
try:
|
||||
await _persist_routstr_token_and_mark_sent(
|
||||
row,
|
||||
operation_id,
|
||||
expected_sats=expected_sats,
|
||||
token=token,
|
||||
amount=amount,
|
||||
mint_url=actual_mint_url,
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Aborting auto top-up because its token and sent claim "
|
||||
"could not be persisted atomically",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
try:
|
||||
await release_token_reservation(token)
|
||||
except Exception as error:
|
||||
logger.critical(
|
||||
"Failed to release untracked auto-topup token",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"mint_url": actual_mint_url,
|
||||
"error": repr(error),
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Auto-topup token was released after persistence failed",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to create cashu token for auto top-up",
|
||||
logger.warning(
|
||||
"Failed to create or persist cashu token for auto top-up",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"amount": amount,
|
||||
"mint_url": mint_url,
|
||||
"error": str(e),
|
||||
"error": repr(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
await _release_routstr_claim(row, operation_id)
|
||||
return
|
||||
|
||||
actual_mint_url = token_mint_url(token, mint_url)
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=actual_mint_url,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Aborting auto top-up because its cashu token could not be persisted",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
try:
|
||||
await release_token_reservation(token)
|
||||
except Exception as error:
|
||||
logger.critical(
|
||||
"Failed to release untracked auto-topup token",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"mint_url": actual_mint_url,
|
||||
"error": str(error),
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Auto-topup token was released after persistence failed",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
await _release_routstr_claim(row, operation_id)
|
||||
return
|
||||
|
||||
# Move the claim before the network call, not after: a worker that dies
|
||||
# mid-request must leave behind a claim that says a token may already be
|
||||
# with the peer.
|
||||
await _mark_routstr_sent(
|
||||
row,
|
||||
operation_id,
|
||||
expected_sats=expected_sats,
|
||||
token=token,
|
||||
amount=amount,
|
||||
mint_url=actual_mint_url,
|
||||
)
|
||||
|
||||
# The audit row and SENT claim committed together before this network call,
|
||||
# so a worker crash cannot make reconciliation treat reserved proofs as an
|
||||
# unspent CLAIMED attempt.
|
||||
result = await provider.topup(token)
|
||||
|
||||
if "error" in result:
|
||||
@@ -705,7 +704,7 @@ async def _release_routstr_claim(row: UpstreamProviderRow, operation_id: str) ->
|
||||
)
|
||||
|
||||
|
||||
async def _mark_routstr_sent(
|
||||
async def _persist_routstr_token_and_mark_sent(
|
||||
row: UpstreamProviderRow,
|
||||
operation_id: str,
|
||||
*,
|
||||
@@ -714,20 +713,60 @@ async def _mark_routstr_sent(
|
||||
amount: int,
|
||||
mint_url: str,
|
||||
) -> None:
|
||||
claim = await _current_routstr_claim(row)
|
||||
failures = claim.failures if claim else 0
|
||||
if not await _advance_routstr_claim(
|
||||
row,
|
||||
operation_id,
|
||||
deadline=int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS,
|
||||
phase=ROUTSTR_PHASE_SENT,
|
||||
expected_sats=expected_sats,
|
||||
failures=failures,
|
||||
token=token,
|
||||
amount=amount,
|
||||
mint_url=mint_url,
|
||||
):
|
||||
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||
"""Commit the bearer-token audit row and SENT claim atomically."""
|
||||
state_id = _routstr_state_id(row)
|
||||
async with create_session() as session:
|
||||
state = await session.get(CashuTransaction, state_id)
|
||||
claim = _parse_routstr_request_id(state.request_id if state else None)
|
||||
if (
|
||||
state is None
|
||||
or state.collected
|
||||
or state.swept
|
||||
or claim is None
|
||||
or claim.operation_id != operation_id
|
||||
or claim.phase != ROUTSTR_PHASE_CLAIMED
|
||||
):
|
||||
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(CashuTransaction)
|
||||
.where(
|
||||
col(CashuTransaction.id) == state_id,
|
||||
col(CashuTransaction.request_id) == state.request_id,
|
||||
col(CashuTransaction.collected) == False, # noqa: E712
|
||||
col(CashuTransaction.swept) == False, # noqa: E712
|
||||
)
|
||||
.values(
|
||||
request_id=_routstr_request_id(
|
||||
operation_id,
|
||||
int(time.time()) + ROUTSTR_PENDING_TTL_SECONDS,
|
||||
ROUTSTR_PHASE_SENT,
|
||||
expected_sats,
|
||||
claim.failures,
|
||||
),
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=mint_url,
|
||||
)
|
||||
)
|
||||
if (getattr(result, "rowcount", 0) or 0) != 1:
|
||||
await session.rollback()
|
||||
raise RuntimeError("Routstr auto top-up claim ownership was lost")
|
||||
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
id=uuid.uuid4().hex,
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=mint_url,
|
||||
type="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def _current_routstr_claim(row: UpstreamProviderRow) -> RoutstrClaim | None:
|
||||
@@ -1081,7 +1120,13 @@ async def _set_ppq_state_terminal(
|
||||
col(CashuTransaction.collected) == False, # noqa: E712
|
||||
col(CashuTransaction.swept) == False, # noqa: E712
|
||||
)
|
||||
.values(collected=collected, swept=swept)
|
||||
.values(
|
||||
collected=collected,
|
||||
swept=swept,
|
||||
# For successful payments this timestamps the durable cooldown,
|
||||
# not merely when the original claim was created.
|
||||
created_at=int(time.time()) if collected else CashuTransaction.created_at,
|
||||
)
|
||||
)
|
||||
updated = (getattr(result, "rowcount", 0) or 0) == 1
|
||||
if updated:
|
||||
@@ -1106,8 +1151,10 @@ async def _reconcile_ppq_state(
|
||||
"""
|
||||
async with create_session() as session:
|
||||
transaction = await session.get(CashuTransaction, _ppq_state_id(row))
|
||||
if transaction is None or transaction.collected or transaction.swept:
|
||||
if transaction is None or transaction.swept:
|
||||
return False
|
||||
if transaction.collected:
|
||||
return int(time.time()) - transaction.created_at < PPQ_SETTLED_COOLDOWN_SECONDS
|
||||
|
||||
claim = _parse_ppq_request_id(transaction.request_id)
|
||||
if claim is None:
|
||||
@@ -1173,7 +1220,7 @@ async def _reconcile_ppq_state(
|
||||
|
||||
|
||||
async def _ppq_provider_is_claimable(
|
||||
session: AsyncSession, provider_id: int | None
|
||||
session: AsyncSession, row: UpstreamProviderRow
|
||||
) -> bool:
|
||||
"""Re-read the provider inside the claim transaction.
|
||||
|
||||
@@ -1184,10 +1231,16 @@ async def _ppq_provider_is_claimable(
|
||||
this the worker could create a claim for a provider that no longer
|
||||
exists, orphaning it forever.
|
||||
"""
|
||||
if provider_id is None:
|
||||
if row.id is None:
|
||||
return False
|
||||
current = await session.get(UpstreamProviderRow, provider_id)
|
||||
return current is not None and current.provider_type == "ppqai"
|
||||
current = await session.get(UpstreamProviderRow, row.id)
|
||||
return bool(
|
||||
current is not None
|
||||
and current.enabled
|
||||
and current.provider_type == "ppqai"
|
||||
and current.api_key == row.api_key
|
||||
and current.provider_settings == row.provider_settings
|
||||
)
|
||||
|
||||
|
||||
async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||
@@ -1198,10 +1251,17 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||
request_id = _ppq_request_id(operation_id, expires_at, PPQ_PHASE_CLAIMED, "pending")
|
||||
|
||||
async with create_session() as session:
|
||||
if not await _ppq_provider_is_claimable(session, row.id):
|
||||
if not await _ppq_provider_is_claimable(session, row):
|
||||
return None
|
||||
existing = await session.get(CashuTransaction, state_id)
|
||||
if existing is not None:
|
||||
if (
|
||||
existing.collected
|
||||
and not existing.swept
|
||||
and int(time.time()) - existing.created_at
|
||||
< PPQ_SETTLED_COOLDOWN_SECONDS
|
||||
):
|
||||
return None
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(CashuTransaction)
|
||||
.where(
|
||||
@@ -1232,7 +1292,7 @@ async def _claim_ppq_topup(row: UpstreamProviderRow) -> str | None:
|
||||
async with create_session() as session:
|
||||
# Same fencing as the update path: the provider must still exist
|
||||
# inside the transaction that creates the claim.
|
||||
if not await _ppq_provider_is_claimable(session, row.id):
|
||||
if not await _ppq_provider_is_claimable(session, row):
|
||||
return None
|
||||
session.add(
|
||||
CashuTransaction(
|
||||
@@ -1453,6 +1513,27 @@ async def _check_and_topup_ppq(row: UpstreamProviderRow, settings: dict) -> None
|
||||
if balance >= threshold_usd:
|
||||
return
|
||||
|
||||
# A single stale/partial balance response must never create an invoice.
|
||||
# Read the uncached endpoint again and require independent agreement.
|
||||
confirmed_balance = await provider.get_balance()
|
||||
if (
|
||||
confirmed_balance is None
|
||||
or not math.isfinite(confirmed_balance)
|
||||
or confirmed_balance < 0
|
||||
or confirmed_balance >= threshold_usd
|
||||
):
|
||||
logger.info(
|
||||
"PPQ auto top-up aborted by balance confirmation",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"first_balance_usd": balance,
|
||||
"confirmed_balance_usd": confirmed_balance,
|
||||
"threshold_usd": threshold_usd,
|
||||
},
|
||||
)
|
||||
return
|
||||
balance = confirmed_balance
|
||||
|
||||
# Perform local pricing and owner-funds checks before asking PPQ to create
|
||||
# an invoice. The exact mint quote still has to be checked afterward, but
|
||||
# predictable local failures should not leave abandoned PPQ invoices.
|
||||
|
||||
+109
-52
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import math
|
||||
import traceback
|
||||
@@ -70,6 +71,17 @@ if typing.TYPE_CHECKING:
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
async def _aclose_if_needed(resource: object | None) -> None:
|
||||
if resource is None:
|
||||
return
|
||||
close = getattr(resource, "aclose", None)
|
||||
if close is None:
|
||||
return
|
||||
result = close()
|
||||
if inspect.isawaitable(result):
|
||||
await result
|
||||
|
||||
|
||||
CostMetadata = CostData | MaxCostData | dict[str, Any]
|
||||
|
||||
|
||||
@@ -993,6 +1005,7 @@ class BaseUpstreamProvider:
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1034,23 +1047,42 @@ class BaseUpstreamProvider:
|
||||
nonlocal usage_finalized
|
||||
if usage_finalized:
|
||||
return
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
try:
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(
|
||||
key.__class__, key.hashed_key
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
pass
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Fallback stream billing finalization failed; releasing reservation",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
usage_finalized = (
|
||||
await self._release_failed_streaming_reservation(
|
||||
fresh_key, new_session, reservation_snapshot
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
# Preserve the original stream exception. If the database
|
||||
# cannot even be opened/read, stale-reservation cleanup is
|
||||
# the only safe recovery path.
|
||||
logger.exception(
|
||||
"Fallback stream billing recovery could not access the database",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
|
||||
def _process_event(
|
||||
raw_event: bytes, final: bool = False
|
||||
@@ -1280,18 +1312,23 @@ class BaseUpstreamProvider:
|
||||
|
||||
except Exception as stream_error:
|
||||
logger.warning(
|
||||
"Streaming interrupted; finalizing in background",
|
||||
"Streaming interrupted; finalizing before closing upstream",
|
||||
extra={
|
||||
"error": str(stream_error),
|
||||
"error_type": type(stream_error).__name__,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
if not usage_finalized:
|
||||
# Create a background task to ensure finalization happens
|
||||
# even if the generator is closed early
|
||||
background_tasks.add_task(finalize_db_only)
|
||||
try:
|
||||
if not usage_finalized:
|
||||
await finalize_db_only()
|
||||
finally:
|
||||
try:
|
||||
await _aclose_if_needed(response)
|
||||
finally:
|
||||
await _aclose_if_needed(client)
|
||||
|
||||
# Remove inaccurate encoding headers from upstream response
|
||||
response_headers = dict(response.headers)
|
||||
@@ -1451,6 +1488,7 @@ class BaseUpstreamProvider:
|
||||
requested_model: str | None = None,
|
||||
model_obj: Model | None = None,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle streaming Responses API responses with token usage tracking and cost adjustment.
|
||||
|
||||
@@ -1484,23 +1522,42 @@ class BaseUpstreamProvider:
|
||||
nonlocal usage_finalized
|
||||
if usage_finalized:
|
||||
return
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
try:
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(
|
||||
key.__class__, key.hashed_key
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
pass
|
||||
if not fresh_key:
|
||||
return
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{"model": last_model_seen or "unknown", "usage": None},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
model_obj,
|
||||
self.provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
usage_finalized = True
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Fallback Responses billing finalization failed; releasing reservation",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
usage_finalized = (
|
||||
await self._release_failed_streaming_reservation(
|
||||
fresh_key, new_session, reservation_snapshot
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
# Preserve the original stream exception. If the database
|
||||
# cannot even be opened/read, stale-reservation cleanup is
|
||||
# the only safe recovery path.
|
||||
logger.exception(
|
||||
"Fallback Responses billing recovery could not access the database",
|
||||
extra={"key_hash": key.hashed_key[:8] + "..."},
|
||||
)
|
||||
|
||||
def _process_event(
|
||||
raw_event: bytes, final: bool = False
|
||||
@@ -1690,16 +1747,23 @@ class BaseUpstreamProvider:
|
||||
|
||||
except Exception as stream_error:
|
||||
logger.warning(
|
||||
"Responses API streaming interrupted; finalizing in background",
|
||||
"Responses API streaming interrupted; finalizing before closing upstream",
|
||||
extra={
|
||||
"error": str(stream_error),
|
||||
"error_type": type(stream_error).__name__,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
if not usage_finalized:
|
||||
await finalize_db_only()
|
||||
try:
|
||||
if not usage_finalized:
|
||||
await finalize_db_only()
|
||||
finally:
|
||||
try:
|
||||
await _aclose_if_needed(response)
|
||||
finally:
|
||||
await _aclose_if_needed(client)
|
||||
|
||||
# Remove inaccurate encoding headers from upstream response
|
||||
response_headers = dict(response.headers)
|
||||
@@ -3051,9 +3115,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
if is_streaming and response.status_code == 200:
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
result = await self.handle_streaming_chat_completion(
|
||||
return await self.handle_streaming_chat_completion(
|
||||
response,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
@@ -3061,9 +3123,8 @@ class BaseUpstreamProvider:
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
reservation_snapshot=reservation_snapshot,
|
||||
client=client,
|
||||
)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
|
||||
# Handle both non-streaming chat completions and embeddings
|
||||
if response.status_code == 200:
|
||||
@@ -3332,19 +3393,15 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
|
||||
if is_streaming and response.status_code == 200:
|
||||
result = await self.handle_streaming_responses_completion(
|
||||
return await self.handle_streaming_responses_completion(
|
||||
response,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
requested_model=original_model_id,
|
||||
model_obj=model_obj,
|
||||
reservation_snapshot=reservation_snapshot,
|
||||
client=client,
|
||||
)
|
||||
background_tasks = BackgroundTasks()
|
||||
background_tasks.add_task(response.aclose)
|
||||
background_tasks.add_task(client.aclose)
|
||||
result.background = background_tasks
|
||||
return result
|
||||
|
||||
if response.status_code == 200:
|
||||
try:
|
||||
@@ -5339,7 +5396,7 @@ class BaseUpstreamProvider:
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Failed to refresh models cache for {self.provider_type or self.base_url}",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
extra={"error": repr(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
def get_cached_models(self) -> list[Model]:
|
||||
|
||||
+101
-12
@@ -1,5 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import httpx
|
||||
@@ -15,6 +19,90 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
_PPQ_SAFE_READ_ATTEMPTS = 3
|
||||
_PPQ_CIRCUIT_COOLDOWN_SECONDS = 30.0
|
||||
|
||||
|
||||
class PPQCircuitOpenError(RuntimeError):
|
||||
"""PPQ safe reads are suppressed until one probe is allowed."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PPQCircuitState:
|
||||
consecutive_failures: int = 0
|
||||
cooldown_until: float = 0.0
|
||||
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
loop: asyncio.AbstractEventLoop | None = None
|
||||
|
||||
|
||||
_ppq_circuits: dict[str, _PPQCircuitState] = {}
|
||||
|
||||
|
||||
def _ppq_origin(url: str) -> str:
|
||||
parsed = httpx.URL(url)
|
||||
return f"{parsed.scheme}://{parsed.host}:{parsed.port}"
|
||||
|
||||
|
||||
async def _safe_read_request(
|
||||
client: httpx.AsyncClient,
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
headers: dict[str, str],
|
||||
json: dict[str, object] | None = None,
|
||||
) -> httpx.Response:
|
||||
"""Retry safe reads, then open one process-local circuit per PPQ origin."""
|
||||
state = _ppq_circuits.setdefault(_ppq_origin(url), _PPQCircuitState())
|
||||
loop = asyncio.get_running_loop()
|
||||
if state.loop is not loop:
|
||||
# Runtime uses one long-lived loop; pytest and some embedded hosts do
|
||||
# not. Preserve circuit state while replacing a loop-bound lock.
|
||||
state.lock = asyncio.Lock()
|
||||
state.loop = loop
|
||||
async with state.lock:
|
||||
remaining = state.cooldown_until - time.monotonic()
|
||||
if remaining > 0:
|
||||
raise PPQCircuitOpenError(
|
||||
f"PPQ.AI safe-read circuit is open; retry after {remaining:.2f}s"
|
||||
)
|
||||
|
||||
for attempt in range(1, _PPQ_SAFE_READ_ATTEMPTS + 1):
|
||||
try:
|
||||
response = await client.request(
|
||||
method, url, headers=headers, json=json
|
||||
)
|
||||
response.raise_for_status()
|
||||
state.consecutive_failures = 0
|
||||
state.cooldown_until = 0.0
|
||||
return response
|
||||
except (httpx.TransportError, httpx.HTTPStatusError) as error:
|
||||
retryable_status = isinstance(error, httpx.HTTPStatusError) and (
|
||||
error.response.status_code in {502, 503, 504}
|
||||
)
|
||||
if not isinstance(error, httpx.TransportError) and not retryable_status:
|
||||
raise
|
||||
state.consecutive_failures += 1
|
||||
if attempt >= _PPQ_SAFE_READ_ATTEMPTS:
|
||||
state.cooldown_until = (
|
||||
time.monotonic() + _PPQ_CIRCUIT_COOLDOWN_SECONDS
|
||||
)
|
||||
raise
|
||||
base_delay = 0.25 * (2 ** (attempt - 1))
|
||||
delay = base_delay + random.uniform(0.0, base_delay)
|
||||
logger.warning(
|
||||
"PPQ.AI safe read failed; retrying",
|
||||
extra={
|
||||
"url": url,
|
||||
"attempt": attempt,
|
||||
"max_attempts": _PPQ_SAFE_READ_ATTEMPTS,
|
||||
"backoff_seconds": round(delay, 3),
|
||||
"error": repr(error),
|
||||
"error_type": type(error).__name__,
|
||||
},
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
raise RuntimeError("unreachable")
|
||||
|
||||
|
||||
class PPQAIModelPricing(BaseModel):
|
||||
ui: Optional[dict[str, float]] = None
|
||||
@@ -125,8 +213,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
response = await _safe_read_request(
|
||||
client, "GET", url, headers=headers
|
||||
)
|
||||
data = response.json()
|
||||
|
||||
models_data = data.get("data", [])
|
||||
@@ -233,12 +322,10 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
return models
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error fetching models from PPQ.AI",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
return []
|
||||
except Exception:
|
||||
# The base refresh handler preserves the last good model cache when
|
||||
# fetching raises; [] would look like a valid empty catalog.
|
||||
raise
|
||||
|
||||
async def on_upstream_error_redirect(
|
||||
self, status_code: int, error_message: str
|
||||
@@ -360,8 +447,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
response = await _safe_read_request(
|
||||
client, "GET", url, headers=headers
|
||||
)
|
||||
status_data = response.json()
|
||||
|
||||
is_paid = status_data.get("status") == "Settled"
|
||||
@@ -460,8 +548,9 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
logger.debug("Checking PPQ.AI account balance", extra={"url": url})
|
||||
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.post(url, headers=headers, json={})
|
||||
response.raise_for_status()
|
||||
response = await _safe_read_request(
|
||||
client, "POST", url, headers=headers, json={}
|
||||
)
|
||||
balance_data = response.json()
|
||||
|
||||
logger.debug(
|
||||
|
||||
+38
-4
@@ -524,7 +524,11 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
|
||||
|
||||
|
||||
async def _send_locked(
|
||||
amount: int, unit: str, mint_url: str | None = None
|
||||
amount: int,
|
||||
unit: str,
|
||||
mint_url: str | None = None,
|
||||
*,
|
||||
owner_only: bool = False,
|
||||
) -> tuple[int, str]:
|
||||
effective_mint_url = await find_trusted_mint_with_funds(
|
||||
amount, unit, mint_url, force_reload=True
|
||||
@@ -534,6 +538,12 @@ async def _send_locked(
|
||||
wallet, effective_mint_url, unit, not_reserved=True
|
||||
)
|
||||
proofs_for_mint = sum(proof.amount for proof in proofs)
|
||||
if owner_only:
|
||||
owner_balance = await _owner_balance_for_mint_and_unit(
|
||||
effective_mint_url, unit, proofs_for_mint
|
||||
)
|
||||
if owner_balance < amount:
|
||||
raise ValueError("Owner Cashu balance is insufficient for auto top-up")
|
||||
all_proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
reserved_for_mint = sum(p.amount for p in all_proofs if p.reserved)
|
||||
|
||||
@@ -580,6 +590,14 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str
|
||||
return token
|
||||
|
||||
|
||||
async def send_token_from_owner_locked(
|
||||
amount: int, unit: str, mint_url: str | None = None
|
||||
) -> str:
|
||||
"""Create an owner-funded token while the caller holds the wallet guard."""
|
||||
_, token = await _send_locked(amount, unit, mint_url, owner_only=True)
|
||||
return token
|
||||
|
||||
|
||||
class Bolt11PaymentNotAttempted(Exception):
|
||||
"""The invoice was definitively not paid, so the attempt can be retried.
|
||||
|
||||
@@ -1824,9 +1842,25 @@ async def _credit_balance_locked(
|
||||
)
|
||||
return amount
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"credit_balance: Error during token redemption",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
classification = classify_redemption_error(e)
|
||||
expected_codes = {
|
||||
"cashu_token_already_spent",
|
||||
"cashu_source_mint_unreachable",
|
||||
"cashu_mint_unreachable",
|
||||
"cashu_mint_rate_limited",
|
||||
}
|
||||
log = (
|
||||
logger.info
|
||||
if classification is not None and classification[3] in expected_codes
|
||||
else logger.error
|
||||
)
|
||||
log(
|
||||
"credit_balance: Token redemption failed",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"error_code": classification[3] if classification else None,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from typing import Any, AsyncGenerator, Callable, Dict, List, Optional, Tuple
|
||||
from typing import Any, AsyncGenerator, Callable, Dict, Iterator, List, Optional, Tuple
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -68,6 +68,15 @@ os.environ.pop("ADMIN_PASSWORD", None)
|
||||
|
||||
from routstr.core.db import ApiKey, get_session # noqa: E402
|
||||
from routstr.core.main import app, lifespan # noqa: E402
|
||||
from routstr.mint import MintRateGuard # noqa: E402
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolate_mint_rate_guards() -> Iterator[None]:
|
||||
"""Do not let one integration test's simulated outage poison the next."""
|
||||
MintRateGuard._guards.clear()
|
||||
yield
|
||||
MintRateGuard._guards.clear()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
|
||||
@@ -23,6 +23,7 @@ from routstr.upstream.auto_topup import (
|
||||
_ppq_request_id,
|
||||
_ppq_spent_last_24h_usd,
|
||||
_ppq_state_id_for_provider,
|
||||
_reconcile_ppq_state,
|
||||
_record_ppq_invoice,
|
||||
_set_ppq_state_terminal,
|
||||
get_ppq_auto_topup_state,
|
||||
@@ -35,6 +36,8 @@ pytestmark = pytest.mark.asyncio
|
||||
def _row(provider_id: int = 1) -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.id = provider_id
|
||||
row.api_key = "secret"
|
||||
row.provider_settings = None
|
||||
return row
|
||||
|
||||
|
||||
@@ -111,12 +114,35 @@ async def test_claim_is_reusable_once_the_previous_attempt_finished(
|
||||
await _seed_provider()
|
||||
first = await _claim_ppq_topup(_row())
|
||||
assert first is not None
|
||||
assert await _set_ppq_state_terminal(_row(), first, collected=True, swept=False)
|
||||
assert await _set_ppq_state_terminal(_row(), first, collected=False, swept=True)
|
||||
|
||||
second = await _claim_ppq_topup(_row())
|
||||
assert second is not None and second != first
|
||||
|
||||
|
||||
async def test_settled_claim_suppresses_immediate_duplicate(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
await _seed_provider()
|
||||
row = _row()
|
||||
operation_id = await _claim_ppq_topup(row)
|
||||
assert operation_id is not None
|
||||
assert await _set_ppq_state_terminal(row, operation_id, collected=True, swept=False)
|
||||
|
||||
assert await _reconcile_ppq_state(row, provider=None) is True
|
||||
assert await _claim_ppq_topup(row) is None
|
||||
|
||||
|
||||
async def test_claim_rejects_stale_provider_configuration(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
await _seed_provider()
|
||||
stale = _row()
|
||||
stale.provider_settings = '{"auto_topup":true}'
|
||||
|
||||
assert await _claim_ppq_topup(stale) is None
|
||||
|
||||
|
||||
async def test_recording_the_invoice_moves_the_claim_in_flight(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
@@ -287,7 +313,15 @@ async def test_ppq_payment_audit_row_is_visible_and_survives_next_claim(
|
||||
assert audit["collected"] is True
|
||||
assert "lnbc-secret-invoice" not in audit["token"]
|
||||
|
||||
# Reusing the deterministic claim lock must not overwrite history.
|
||||
# The durable cooldown blocks an immediate duplicate, then the claim lock
|
||||
# can be reused without overwriting audit history after it expires.
|
||||
assert await _claim_ppq_topup(_row()) is None
|
||||
async with create_session() as session:
|
||||
state = await session.get(CashuTransaction, _ppq_state_id_for_provider(1))
|
||||
assert state is not None
|
||||
state.created_at = int(time.time()) - 301
|
||||
session.add(state)
|
||||
await session.commit()
|
||||
assert await _claim_ppq_topup(_row()) is not None
|
||||
async with create_session() as session:
|
||||
assert await session.get(CashuTransaction, audit["id"]) is not None
|
||||
|
||||
@@ -25,6 +25,7 @@ from routstr.upstream.auto_topup import (
|
||||
_check_and_topup,
|
||||
_claim_routstr_topup,
|
||||
_parse_routstr_request_id,
|
||||
_persist_routstr_token_and_mark_sent,
|
||||
_routstr_spent_last_24h_sats,
|
||||
_routstr_state_id_for_provider,
|
||||
get_routstr_auto_topup_state,
|
||||
@@ -75,7 +76,9 @@ def _patch_wallet(module: Any, peer: MagicMock, token: str) -> ExitStack:
|
||||
patch.object(module.RoutstrUpstreamProvider, "from_db_row", return_value=peer)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(module, "send_token", AsyncMock(return_value=token))
|
||||
patch.object(
|
||||
module, "send_token_from_owner_locked", AsyncMock(return_value=token)
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch.object(module, "token_mint_url", return_value="https://mint.test")
|
||||
@@ -111,6 +114,32 @@ async def test_second_worker_cannot_claim_while_the_first_holds_one(
|
||||
assert await _claim_routstr_topup(row, expected_sats=TOPUP_SATS) is None
|
||||
|
||||
|
||||
async def test_token_and_sent_claim_roll_back_together_on_commit_failure(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
row = await _seed_provider()
|
||||
operation_id = await _claim_routstr_topup(row, expected_sats=TOPUP_SATS)
|
||||
assert operation_id is not None
|
||||
|
||||
with patch(
|
||||
"sqlmodel.ext.asyncio.session.AsyncSession.commit",
|
||||
new=AsyncMock(side_effect=RuntimeError("commit failed")),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="commit failed"):
|
||||
await _persist_routstr_token_and_mark_sent(
|
||||
row,
|
||||
operation_id,
|
||||
expected_sats=TOPUP_SATS,
|
||||
token="cashu-token-atomic",
|
||||
amount=TOPUP_SATS,
|
||||
mint_url="https://mint.test",
|
||||
)
|
||||
|
||||
claim = _parse_routstr_request_id((await _claim_state()).request_id) # type: ignore[union-attr]
|
||||
assert claim is not None and claim.phase != ROUTSTR_PHASE_SENT
|
||||
assert await _sent_tokens() == []
|
||||
|
||||
|
||||
async def test_token_is_persisted_before_it_reaches_the_peer(
|
||||
patched_db_engine: Any,
|
||||
) -> None:
|
||||
@@ -142,7 +171,7 @@ async def test_untracked_token_is_returned_and_never_sent(
|
||||
_patch_wallet(auto_topup_module, peer, "cashu-token-1"),
|
||||
patch.object(
|
||||
auto_topup_module,
|
||||
"store_cashu_transaction",
|
||||
"_persist_routstr_token_and_mark_sent",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
),
|
||||
patch.object(
|
||||
|
||||
@@ -3,6 +3,7 @@ Integration tests for wallet authentication system including API key generation
|
||||
Tests POST /v1/wallet/topup endpoint and authorization header validation.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any
|
||||
|
||||
@@ -113,38 +114,35 @@ async def test_api_key_generation_invalid_token(
|
||||
async def test_duplicate_token_handling(
|
||||
integration_client: AsyncClient, testmint_wallet: Any, db_snapshot: Any
|
||||
) -> None:
|
||||
"""Test that duplicate tokens return the same API key without double-spending"""
|
||||
"""Concurrent duplicate redemption credits once and deterministically replays."""
|
||||
|
||||
# Generate a valid token
|
||||
amount = 500 # 500 sats
|
||||
amount = 500
|
||||
token = await testmint_wallet.mint_tokens(amount)
|
||||
|
||||
# First use of token
|
||||
integration_client.headers["Authorization"] = f"Bearer {token}"
|
||||
response1 = await integration_client.get("/v1/wallet/info")
|
||||
assert response1.status_code == 200
|
||||
|
||||
response1, response2 = await asyncio.gather(
|
||||
integration_client.get("/v1/wallet/info"),
|
||||
integration_client.get("/v1/wallet/info"),
|
||||
)
|
||||
assert response1.status_code < 500
|
||||
assert response2.status_code < 500
|
||||
assert response1.status_code == response2.status_code == 200
|
||||
api_key1 = response1.json()["api_key"]
|
||||
balance1 = response1.json()["balance"]
|
||||
|
||||
# Capture state after first submission
|
||||
await db_snapshot.capture()
|
||||
|
||||
# Second use of same token - should return same API key since it's already created
|
||||
response2 = await integration_client.get("/v1/wallet/info")
|
||||
assert response2.status_code == 200
|
||||
api_key2 = response2.json()["api_key"]
|
||||
balance1 = response1.json()["balance"]
|
||||
balance2 = response2.json()["balance"]
|
||||
|
||||
# Should return the same API key and balance
|
||||
assert api_key1 == api_key2
|
||||
assert balance1 == balance2
|
||||
assert balance1 == balance2 == amount * 1000
|
||||
|
||||
# Verify no additional database changes
|
||||
# The real DB contains one logical credit. A later replay is also mutation-free.
|
||||
await db_snapshot.capture()
|
||||
replay = await integration_client.get("/v1/wallet/info")
|
||||
assert replay.status_code == 200
|
||||
assert replay.json()["api_key"] == api_key1
|
||||
diff = await db_snapshot.diff()
|
||||
assert len(diff["api_keys"]["added"]) == 0
|
||||
assert len(diff["api_keys"]["modified"]) == 0
|
||||
|
||||
# Original API key should still work with original balance
|
||||
integration_client.headers["Authorization"] = f"Bearer {api_key1}"
|
||||
wallet_response = await integration_client.get("/v1/wallet/")
|
||||
assert wallet_response.status_code == 200
|
||||
|
||||
@@ -10,9 +10,12 @@ invalidate them on "paid" or release them on "unpaid".
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
from cashu.core.base import MeltQuote, MeltQuoteState, Proof
|
||||
from cashu.core.models import PostMeltQuoteResponse
|
||||
from cashu.wallet import crud
|
||||
from cashu.wallet.wallet import Wallet
|
||||
|
||||
@@ -47,6 +50,20 @@ async def _seed_ambiguous_melt(wallet: Wallet) -> list[Proof]:
|
||||
proofs = [_proof("secret-a"), _proof("secret-b", amount=32)]
|
||||
for proof in proofs:
|
||||
await crud.store_proof(proof, db=wallet.db)
|
||||
await crud.store_bolt11_melt_quote(
|
||||
db=wallet.db,
|
||||
quote=MeltQuote(
|
||||
quote=QUOTE_ID,
|
||||
method="bolt11",
|
||||
request="lnbc1-test",
|
||||
checking_id="",
|
||||
unit="sat",
|
||||
amount=95,
|
||||
fee_reserve=1,
|
||||
state=MeltQuoteState.pending,
|
||||
mint=str(wallet.url),
|
||||
),
|
||||
)
|
||||
|
||||
await wallet.set_reserved_for_melt(proofs, reserved=True, quote_id=QUOTE_ID)
|
||||
# cashu's `except` block in melt():
|
||||
@@ -75,6 +92,35 @@ async def test_melt_recovery_is_findable_by_quote_after_restart(
|
||||
assert all(p.melt_id == QUOTE_ID for p in found)
|
||||
|
||||
|
||||
async def test_paid_reconciliation_invalidates_recovered_proofs_after_restart(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Actual Wallet.get_melt_quote consumes quote-linked proofs after restart."""
|
||||
wallet = await _wallet(tmp_path)
|
||||
await _seed_ambiguous_melt(wallet)
|
||||
|
||||
restarted = await _wallet(tmp_path)
|
||||
remote = PostMeltQuoteResponse(
|
||||
quote=QUOTE_ID,
|
||||
amount=95,
|
||||
unit="sat",
|
||||
request="lnbc1-test",
|
||||
fee_reserve=1,
|
||||
state=MeltQuoteState.paid.value,
|
||||
expiry=None,
|
||||
payment_preimage="preimage",
|
||||
)
|
||||
with patch(
|
||||
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
|
||||
new=AsyncMock(return_value=remote),
|
||||
):
|
||||
reconciled = await restarted.get_melt_quote(QUOTE_ID)
|
||||
|
||||
assert reconciled is not None and reconciled.state == MeltQuoteState.paid
|
||||
assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == []
|
||||
assert await crud.get_proofs(db=restarted.db) == []
|
||||
|
||||
|
||||
async def test_send_style_reservation_would_not_be_reconcilable(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
@@ -97,19 +143,65 @@ async def test_send_style_reservation_would_not_be_reconcilable(
|
||||
async def test_unpaid_reconciliation_releases_recovered_proofs_after_restart(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""The full recovery arc: crash, restart, mint says unpaid, funds usable."""
|
||||
"""Actual Wallet.get_melt_quote releases proofs after an unpaid answer."""
|
||||
wallet = await _wallet(tmp_path)
|
||||
await _seed_ambiguous_melt(wallet)
|
||||
|
||||
restarted = await _wallet(tmp_path)
|
||||
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert len(found) == 2
|
||||
remote = PostMeltQuoteResponse(
|
||||
quote=QUOTE_ID,
|
||||
amount=95,
|
||||
unit="sat",
|
||||
request="lnbc1-test",
|
||||
fee_reserve=1,
|
||||
state=MeltQuoteState.unpaid.value,
|
||||
expiry=None,
|
||||
)
|
||||
with patch(
|
||||
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
|
||||
new=AsyncMock(return_value=remote),
|
||||
):
|
||||
reconciled = await restarted.get_melt_quote(QUOTE_ID)
|
||||
|
||||
# What get_melt_quote() does on an "unpaid" answer.
|
||||
await restarted.set_reserved_for_melt(found, reserved=False, quote_id=None)
|
||||
|
||||
released = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert released == []
|
||||
assert reconciled is not None and reconciled.state == MeltQuoteState.unpaid
|
||||
assert await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID) == []
|
||||
all_proofs = await crud.get_proofs(db=restarted.db)
|
||||
assert len(all_proofs) == 2
|
||||
assert all(not p.reserved for p in all_proofs) # spendable again
|
||||
assert all(not p.reserved for p in all_proofs)
|
||||
|
||||
|
||||
async def test_pending_and_transport_reconciliation_keep_recovered_reservation(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
wallet = await _wallet(tmp_path)
|
||||
await _seed_ambiguous_melt(wallet)
|
||||
restarted = await _wallet(tmp_path)
|
||||
pending = PostMeltQuoteResponse(
|
||||
quote=QUOTE_ID,
|
||||
amount=95,
|
||||
unit="sat",
|
||||
request="lnbc1-test",
|
||||
fee_reserve=1,
|
||||
state=MeltQuoteState.pending.value,
|
||||
expiry=None,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
|
||||
new=AsyncMock(return_value=pending),
|
||||
):
|
||||
reconciled = await restarted.get_melt_quote(QUOTE_ID)
|
||||
assert reconciled is not None and reconciled.state == MeltQuoteState.pending
|
||||
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert len(found) == 2 and all(proof.reserved for proof in found)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"cashu.wallet.v1_api.LedgerAPI.get_melt_quote",
|
||||
new=AsyncMock(side_effect=httpx.ReadTimeout("mint unavailable")),
|
||||
),
|
||||
pytest.raises(httpx.ReadTimeout),
|
||||
):
|
||||
await restarted.get_melt_quote(QUOTE_ID)
|
||||
found = await crud.get_proofs(db=restarted.db, melt_id=QUOTE_ID)
|
||||
assert len(found) == 2 and all(proof.reserved for proof in found)
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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,29 @@ LNURL_DATA = {
|
||||
"max_sendable": 100_000_000,
|
||||
}
|
||||
|
||||
# 1000 sat minus the 11 sat estimated fee reserve.
|
||||
EXPECTED_QUOTE_SAT = 989
|
||||
# The exact plan spends the 1000 sat gross budget as 999 sat + 1 sat reserve.
|
||||
EXPECTED_QUOTE_SAT = 999
|
||||
|
||||
|
||||
def _wallet(
|
||||
quote_amount: int = EXPECTED_QUOTE_SAT,
|
||||
quote_amount: int | None = None,
|
||||
) -> tuple[MagicMock, list[MagicMock]]:
|
||||
proofs = [MagicMock(amount=1000)]
|
||||
proofs = [MagicMock(amount=1000, reserved=False)]
|
||||
wallet = MagicMock(url="https://mint.test")
|
||||
wallet.melt_quote = AsyncMock(
|
||||
return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount)
|
||||
)
|
||||
wallet.get_fees_for_proofs.return_value = 0
|
||||
if quote_amount is None:
|
||||
wallet.melt_quote = AsyncMock(
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=1, quote="q", amount=1000),
|
||||
MagicMock(fee_reserve=1, quote="q", amount=EXPECTED_QUOTE_SAT),
|
||||
]
|
||||
)
|
||||
else:
|
||||
wallet.melt_quote = AsyncMock(
|
||||
return_value=MagicMock(fee_reserve=1, quote="q", amount=quote_amount)
|
||||
)
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
wallet.get_fees_for_proofs = MagicMock(return_value=0)
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
return wallet, proofs
|
||||
@@ -124,16 +135,22 @@ async def test_raw_send_to_lnurl_accepts_exact_invoice() -> None:
|
||||
)
|
||||
|
||||
assert paid == EXPECTED_QUOTE_SAT * 1000
|
||||
wallet.select_to_send.assert_awaited_once()
|
||||
wallet.select_to_send.assert_not_awaited()
|
||||
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=True)
|
||||
wallet.melt.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None:
|
||||
# 1_000_000 msat minus an 11 sat fee reserve leaves 989_000 msat.
|
||||
wallet, proofs = _wallet(989_000)
|
||||
# 1_000_000 msat gross minus the exact 1 msat reserve leaves 999_999 msat.
|
||||
wallet, proofs = _wallet()
|
||||
proofs[0].amount = 1_000_000
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
wallet.melt_quote = AsyncMock(
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=1, quote="q", amount=1_000_000),
|
||||
MagicMock(fee_reserve=1, quote="q", amount=999_999),
|
||||
]
|
||||
)
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with data_patch, invoice_patch:
|
||||
@@ -141,7 +158,46 @@ async def test_raw_send_to_lnurl_msat_unit_compares_in_wallet_unit() -> None:
|
||||
wallet, proofs, "owner@ln.tld", "msat", amount=1_000_000
|
||||
)
|
||||
|
||||
assert paid == 989_000
|
||||
assert paid == 999_999
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_requotes_for_exact_input_fees_without_recursion() -> (
|
||||
None
|
||||
):
|
||||
proofs = [MagicMock(amount=1, reserved=False) for _ in range(1500)]
|
||||
wallet = MagicMock(url="https://mint.test")
|
||||
wallet.get_fees_for_proofs = MagicMock(
|
||||
side_effect=lambda selected: math.ceil(len(selected) / 100)
|
||||
)
|
||||
wallet.melt_quote = AsyncMock(
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=10, quote="q1", amount=1500),
|
||||
MagicMock(fee_reserve=10, quote="q2", amount=1475),
|
||||
]
|
||||
)
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.paid))
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
checkpoint = AsyncMock()
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with data_patch, invoice_patch:
|
||||
paid = await raw_send_to_lnurl(
|
||||
wallet,
|
||||
proofs,
|
||||
"owner@ln.tld",
|
||||
"sat",
|
||||
amount=1500,
|
||||
on_melt_quote=checkpoint,
|
||||
)
|
||||
|
||||
assert paid == 1_475_000
|
||||
assert wallet.melt_quote.await_count == 2
|
||||
checkpoint.assert_awaited_once_with("q2")
|
||||
wallet.select_to_send.assert_not_called()
|
||||
selected = wallet.melt.await_args.kwargs["proofs"]
|
||||
assert sum(proof.amount for proof in selected) == 1500
|
||||
assert 1475 + 10 + wallet.get_fees_for_proofs(selected) == 1500
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""LNURL melt attempts must not misclassify ambiguous payment outcomes."""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
@@ -16,6 +17,14 @@ from routstr.payment.lnurl import (
|
||||
raw_send_to_lnurl,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_mint_guards() -> Iterator[None]:
|
||||
MintRateGuard._guards.clear()
|
||||
yield
|
||||
MintRateGuard._guards.clear()
|
||||
|
||||
|
||||
LNURL_DATA = {
|
||||
"callback_url": "https://ln.tld/cb",
|
||||
"min_sendable": 1_000,
|
||||
@@ -23,18 +32,22 @@ LNURL_DATA = {
|
||||
}
|
||||
|
||||
|
||||
# 1000 sat minus the 11 sat estimated fee reserve.
|
||||
QUOTE_AMOUNT_SAT = 989
|
||||
# Exact planning pays 999 sat from a 1000 sat budget with a 1 sat reserve.
|
||||
QUOTE_AMOUNT_SAT = 999
|
||||
|
||||
|
||||
def _wallet() -> tuple[MagicMock, list[MagicMock]]:
|
||||
proofs = [MagicMock(amount=1000)]
|
||||
proofs = [MagicMock(amount=1000, reserved=False)]
|
||||
wallet = MagicMock(url="https://mint.test")
|
||||
wallet.melt_quote = AsyncMock(
|
||||
return_value=MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT)
|
||||
side_effect=[
|
||||
MagicMock(fee_reserve=1, quote="q", amount=1000),
|
||||
MagicMock(fee_reserve=1, quote="q", amount=QUOTE_AMOUNT_SAT),
|
||||
]
|
||||
)
|
||||
wallet.melt = AsyncMock()
|
||||
wallet.select_to_send = AsyncMock(return_value=(proofs, None))
|
||||
wallet.get_fees_for_proofs = MagicMock(return_value=0)
|
||||
wallet.set_reserved_for_melt = AsyncMock()
|
||||
wallet.set_reserved_for_send = AsyncMock()
|
||||
return wallet, proofs
|
||||
@@ -54,7 +67,26 @@ def _lnurl_patches() -> tuple[Any, Any]:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_timeout_reconciled_unpaid_is_retry_safe() -> None:
|
||||
async def test_raw_send_to_lnurl_direct_unpaid_is_retry_safe() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.unpaid))
|
||||
wallet.get_melt_quote = AsyncMock()
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised,
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
assert not isinstance(raised.value, MeltOutcomeAmbiguousError)
|
||||
wallet.get_melt_quote.assert_not_awaited()
|
||||
wallet.set_reserved_for_send.assert_any_await(proofs, reserved=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_timeout_then_unpaid_remains_ambiguous() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
|
||||
async def _hang(**kwargs: object) -> None:
|
||||
@@ -71,19 +103,19 @@ async def test_raw_send_to_lnurl_timeout_reconciled_unpaid_is_retry_safe() -> No
|
||||
patch.object(settings, "mint_retry_max_attempts", 0),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised,
|
||||
pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"),
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
assert not isinstance(raised.value, MeltOutcomeAmbiguousError)
|
||||
wallet.get_melt_quote.assert_awaited_once_with("q")
|
||||
wallet.set_reserved_for_melt.assert_awaited_once_with(
|
||||
assert wallet.set_reserved_for_melt.await_count == 2
|
||||
wallet.set_reserved_for_melt.assert_awaited_with(
|
||||
proofs, reserved=True, quote_id="q"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_send_to_lnurl_wrapped_transport_error_is_reconciled_once() -> None:
|
||||
async def test_raw_send_to_lnurl_wrapped_transport_unpaid_remains_ambiguous() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
|
||||
async def _wrapped_transport_error(**kwargs: object) -> None:
|
||||
@@ -102,14 +134,14 @@ async def test_raw_send_to_lnurl_wrapped_transport_error_is_reconciled_once() ->
|
||||
patch.object(settings, "mint_retry_max_attempts", 3),
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(LNURLError, match="confirmed that the melt was unpaid") as raised,
|
||||
pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"),
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
assert not isinstance(raised.value, MeltOutcomeAmbiguousError)
|
||||
wallet.melt.assert_awaited_once()
|
||||
wallet.get_melt_quote.assert_awaited_once_with("q")
|
||||
wallet.set_reserved_for_melt.assert_awaited_once_with(
|
||||
assert wallet.set_reserved_for_melt.await_count == 2
|
||||
wallet.set_reserved_for_melt.assert_awaited_with(
|
||||
proofs, reserved=True, quote_id="q"
|
||||
)
|
||||
|
||||
@@ -180,6 +212,45 @@ async def test_raw_send_to_lnurl_pending_response_stays_ambiguous() -> None:
|
||||
wallet.get_melt_quote.assert_awaited_once_with("q")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pending_then_immediate_unpaid_remains_reserved_and_ambiguous() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending))
|
||||
wallet.get_melt_quote = AsyncMock(
|
||||
return_value=MagicMock(state=MeltQuoteState.unpaid)
|
||||
)
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(MeltOutcomeAmbiguousError, match="immediate unpaid"),
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
wallet.set_reserved_for_melt.assert_awaited_once_with(
|
||||
proofs, reserved=True, quote_id="q"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_immediate_unpaid_reservation_failure_stays_ambiguous() -> None:
|
||||
wallet, proofs = _wallet()
|
||||
wallet.melt = AsyncMock(return_value=MagicMock(state=MeltQuoteState.pending))
|
||||
wallet.get_melt_quote = AsyncMock(
|
||||
return_value=MagicMock(state=MeltQuoteState.unpaid)
|
||||
)
|
||||
wallet.set_reserved_for_melt = AsyncMock(side_effect=OSError("db locked"))
|
||||
data_patch, invoice_patch = _lnurl_patches()
|
||||
|
||||
with (
|
||||
data_patch,
|
||||
invoice_patch,
|
||||
pytest.raises(MeltOutcomeAmbiguousError, match="could not be restored"),
|
||||
):
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("rate_error", ["cooldown", "http_429"])
|
||||
async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs(
|
||||
@@ -213,7 +284,8 @@ async def test_raw_send_to_lnurl_rate_rejection_unreserves_proofs(
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
wallet.melt.assert_not_awaited()
|
||||
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=False)
|
||||
assert wallet.set_reserved_for_send.await_count == 2
|
||||
wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -238,7 +310,8 @@ async def test_real_mint_wrapper_http_429_unreserves_proofs() -> None:
|
||||
await raw_send_to_lnurl(wallet, proofs, "owner@ln.tld", "sat", amount=1000)
|
||||
|
||||
wallet.melt.assert_awaited_once()
|
||||
wallet.set_reserved_for_send.assert_awaited_once_with(proofs, reserved=False)
|
||||
assert wallet.set_reserved_for_send.await_count == 2
|
||||
wallet.set_reserved_for_send.assert_awaited_with(proofs, reserved=False)
|
||||
MintRateGuard._guards.pop(str(wallet.url), None)
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from routstr.mint import (
|
||||
MintRateGuard,
|
||||
MintRateLimitedError,
|
||||
fail_fast_mint_operations,
|
||||
run_mint_operation,
|
||||
)
|
||||
from routstr.wallet import Wallet
|
||||
|
||||
@@ -103,6 +104,34 @@ async def test_cashu_429_dispatches_through_wallet_override() -> None:
|
||||
await wallet.mint_quote(1, Unit.sat)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapped_transport_failure_opens_central_cooldown() -> None:
|
||||
mint_url = "https://transport-failure.test"
|
||||
MintRateGuard._guards.pop(mint_url, None)
|
||||
|
||||
async def wrapped_failure() -> None:
|
||||
try:
|
||||
raise httpx.ReadTimeout("body stalled")
|
||||
except httpx.ReadTimeout as error:
|
||||
raise Exception("wallet wrapper") from error
|
||||
|
||||
with pytest.raises(Exception, match="wallet wrapper"):
|
||||
await run_mint_operation(
|
||||
wrapped_failure,
|
||||
mint_url=mint_url,
|
||||
retry_timeouts=False,
|
||||
)
|
||||
|
||||
guard = MintRateGuard.get(mint_url)
|
||||
assert guard.cooldown_remaining() > 29
|
||||
probe = AsyncMock()
|
||||
async with fail_fast_mint_operations():
|
||||
with pytest.raises(MintCooldownError):
|
||||
await guard.run(probe)
|
||||
probe.assert_not_awaited()
|
||||
MintRateGuard._guards.pop(mint_url, None)
|
||||
|
||||
|
||||
async def test_guard_concurrency_change_preserves_cooldown_state() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.ppqai import (
|
||||
PPQAIUpstreamProvider,
|
||||
PPQCircuitOpenError,
|
||||
_ppq_circuits,
|
||||
_safe_read_request,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_ppq_circuits() -> None:
|
||||
_ppq_circuits.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_safe_ppq_read_retries_timeout_with_bounded_backoff() -> None:
|
||||
request = httpx.Request("GET", "https://api.ppq.ai/models")
|
||||
success = httpx.Response(200, request=request, json={"data": []})
|
||||
client = MagicMock()
|
||||
client.request = AsyncMock(
|
||||
side_effect=[httpx.ReadTimeout("", request=request), success]
|
||||
)
|
||||
with (
|
||||
patch("routstr.upstream.ppqai.random.uniform", return_value=0.0),
|
||||
patch("routstr.upstream.ppqai.asyncio.sleep", AsyncMock()) as sleep,
|
||||
):
|
||||
response = await _safe_read_request(
|
||||
client, "GET", "https://api.ppq.ai/models", headers={}
|
||||
)
|
||||
|
||||
assert response is success
|
||||
assert client.request.await_count == 2
|
||||
sleep.assert_awaited_once_with(0.25)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_safe_ppq_read_opens_cross_cycle_circuit_and_probe_clears_it() -> None:
|
||||
request = httpx.Request("GET", "https://api.ppq.ai/models")
|
||||
client = MagicMock()
|
||||
client.request = AsyncMock(side_effect=httpx.ReadTimeout("down", request=request))
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.ppqai.random.uniform", return_value=0.0),
|
||||
patch("routstr.upstream.ppqai.asyncio.sleep", AsyncMock()),
|
||||
patch("routstr.upstream.ppqai.time.monotonic", return_value=100.0),
|
||||
pytest.raises(httpx.ReadTimeout),
|
||||
):
|
||||
await _safe_read_request(client, "GET", str(request.url), headers={})
|
||||
assert client.request.await_count == 3
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.ppqai.time.monotonic", return_value=110.0),
|
||||
pytest.raises(PPQCircuitOpenError),
|
||||
):
|
||||
await _safe_read_request(client, "GET", str(request.url), headers={})
|
||||
assert client.request.await_count == 3
|
||||
|
||||
success = httpx.Response(200, request=request, json={"data": []})
|
||||
client.request = AsyncMock(return_value=success)
|
||||
with patch("routstr.upstream.ppqai.time.monotonic", return_value=131.0):
|
||||
assert (
|
||||
await _safe_read_request(client, "GET", str(request.url), headers={})
|
||||
is success
|
||||
)
|
||||
|
||||
state = next(iter(_ppq_circuits.values()))
|
||||
assert state.consecutive_failures == 0
|
||||
assert state.cooldown_until == 0.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_models_failure_is_not_a_valid_empty_catalog() -> None:
|
||||
provider = PPQAIUpstreamProvider("secret")
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.ppqai._safe_read_request",
|
||||
AsyncMock(side_effect=httpx.ReadTimeout("catalog timed out")),
|
||||
),
|
||||
pytest.raises(httpx.ReadTimeout),
|
||||
):
|
||||
await provider.fetch_models()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ppq_invoice_creation_post_is_never_retried() -> None:
|
||||
provider = PPQAIUpstreamProvider("secret")
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(side_effect=httpx.ReadTimeout("invoice timed out"))
|
||||
context = MagicMock()
|
||||
context.__aenter__ = AsyncMock(return_value=client)
|
||||
context.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.ppqai.httpx.AsyncClient", return_value=context),
|
||||
pytest.raises(httpx.ReadTimeout),
|
||||
):
|
||||
await provider.create_lightning_topup(10, "USD")
|
||||
|
||||
client.post.assert_awaited_once()
|
||||
@@ -4,6 +4,7 @@ from collections.abc import AsyncGenerator
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import BackgroundTasks
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
@@ -340,6 +341,153 @@ async def test_responses_streaming_releases_and_raises_on_billing_failure(
|
||||
release.assert_awaited_once_with(snapshot, session, 500)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("api", ["chat", "responses"])
|
||||
@pytest.mark.parametrize("finalization_fails", [False, True])
|
||||
async def test_partial_remote_protocol_error_finalizes_and_closes_once(
|
||||
api: str,
|
||||
finalization_fails: bool,
|
||||
) -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
|
||||
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
|
||||
yield b'data: {"model":"test","choices":[{"delta":{"content":"hi"}}]}\n\n'
|
||||
raise httpx.RemoteProtocolError("incomplete chunked read")
|
||||
|
||||
upstream_response = MagicMock(
|
||||
status_code=200, headers={"content-type": "text/event-stream"}
|
||||
)
|
||||
upstream_response.aiter_bytes = aiter_bytes
|
||||
upstream_response.aclose = AsyncMock()
|
||||
client = MagicMock()
|
||||
client.aclose = AsyncMock()
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = f"{api}-partial"
|
||||
key.balance = 10_000
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.rollback = AsyncMock()
|
||||
session_context = MagicMock()
|
||||
session_context.__aenter__ = AsyncMock(return_value=session)
|
||||
session_context.__aexit__ = AsyncMock(return_value=None)
|
||||
adjust = (
|
||||
AsyncMock(side_effect=SQLAlchemyError("database unavailable"))
|
||||
if finalization_fails
|
||||
else AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0})
|
||||
)
|
||||
snapshot = ReservationSnapshot(
|
||||
release_id=f"{api}-partial-release",
|
||||
key_hash=key.hashed_key,
|
||||
billing_key_hash=key.hashed_key,
|
||||
reserved_msats=500,
|
||||
)
|
||||
release = AsyncMock(return_value=True)
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.base.adjust_payment_for_tokens", adjust),
|
||||
patch("routstr.upstream.base.release_reservation", release),
|
||||
patch("routstr.upstream.base.create_session", return_value=session_context),
|
||||
):
|
||||
if api == "chat":
|
||||
response = await provider.handle_streaming_chat_completion(
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
background_tasks=BackgroundTasks(),
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
else:
|
||||
response = await provider.handle_streaming_responses_completion(
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
emitted = bytearray()
|
||||
with pytest.raises(httpx.RemoteProtocolError):
|
||||
async for chunk in response.body_iterator:
|
||||
emitted.extend(
|
||||
chunk.encode() if isinstance(chunk, str) else bytes(chunk)
|
||||
)
|
||||
|
||||
adjust.assert_awaited_once()
|
||||
if finalization_fails:
|
||||
session.rollback.assert_awaited_once()
|
||||
release.assert_awaited_once_with(snapshot, session, 500)
|
||||
else:
|
||||
release.assert_not_awaited()
|
||||
upstream_response.aclose.assert_awaited_once()
|
||||
client.aclose.assert_awaited_once()
|
||||
assert b"[DONE]" not in emitted
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("api", ["chat", "responses"])
|
||||
async def test_partial_stream_preserves_transport_error_when_billing_db_is_down(
|
||||
api: str,
|
||||
) -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
|
||||
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
|
||||
yield b'data: {"model":"test","choices":[]}\n\n'
|
||||
raise httpx.RemoteProtocolError("incomplete chunked read")
|
||||
|
||||
upstream_response = MagicMock(
|
||||
status_code=200, headers={"content-type": "text/event-stream"}
|
||||
)
|
||||
upstream_response.aiter_bytes = aiter_bytes
|
||||
upstream_response.aclose = AsyncMock()
|
||||
client = MagicMock()
|
||||
client.aclose = AsyncMock()
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = f"{api}-database-down"
|
||||
key.balance = 10_000
|
||||
snapshot = ReservationSnapshot(
|
||||
release_id=f"{api}-database-down-release",
|
||||
key_hash=key.hashed_key,
|
||||
billing_key_hash=key.hashed_key,
|
||||
reserved_msats=500,
|
||||
)
|
||||
unavailable_session = MagicMock()
|
||||
unavailable_session.__aenter__ = AsyncMock(
|
||||
side_effect=SQLAlchemyError("database unavailable")
|
||||
)
|
||||
unavailable_session.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch(
|
||||
"routstr.upstream.base.create_session", return_value=unavailable_session
|
||||
):
|
||||
if api == "chat":
|
||||
response = await provider.handle_streaming_chat_completion(
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
background_tasks=BackgroundTasks(),
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
else:
|
||||
response = await provider.handle_streaming_responses_completion(
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
with pytest.raises(httpx.RemoteProtocolError, match="incomplete chunked read"):
|
||||
async for _ in response.body_iterator:
|
||||
pass
|
||||
|
||||
upstream_response.aclose.assert_awaited_once()
|
||||
client.aclose.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_streaming_duplicate_publishes_zero_settled_cost() -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
@@ -578,9 +726,7 @@ async def test_client_disconnect_midstream_finalizes_and_stops_heartbeat() -> No
|
||||
background_tasks=background_tasks,
|
||||
reservation_snapshot=snapshot,
|
||||
)
|
||||
iterator = cast(
|
||||
AsyncGenerator[bytes, None], response.body_iterator
|
||||
)
|
||||
iterator = cast(AsyncGenerator[bytes, None], response.body_iterator)
|
||||
await iterator.__anext__() # first chunk reaches the client
|
||||
await iterator.aclose() # client aborts the socket here
|
||||
|
||||
|
||||
@@ -26,6 +26,7 @@ from routstr.wallet import (
|
||||
recieve_token,
|
||||
send,
|
||||
send_token,
|
||||
send_token_from_owner_locked,
|
||||
)
|
||||
|
||||
|
||||
@@ -455,6 +456,33 @@ async def test_send_token() -> None:
|
||||
assert token == "test_token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_owner_only_token_rejects_customer_backed_proofs() -> None:
|
||||
mint = "http://mint:3338"
|
||||
proof = Mock(amount=1000, reserved=False)
|
||||
wallet = Mock(keysets={}, proofs=[proof], select_to_send=AsyncMock())
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.wallet.find_trusted_mint_with_funds",
|
||||
AsyncMock(return_value=mint),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
return_value=[proof],
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet._owner_balance_for_mint_and_unit",
|
||||
AsyncMock(return_value=50),
|
||||
),
|
||||
pytest.raises(ValueError, match="Owner Cashu balance"),
|
||||
):
|
||||
await send_token_from_owner_locked(100, "sat", mint)
|
||||
|
||||
wallet.select_to_send.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_token_reservation_unreserves_local_proofs() -> None:
|
||||
from routstr.wallet import release_token_reservation
|
||||
@@ -695,6 +723,38 @@ async def test_credit_balance() -> None:
|
||||
assert mock_session.refresh.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_duplicate_token_credits_exactly_once() -> None:
|
||||
key = Mock(balance=0, hashed_key="duplicate-key")
|
||||
session = AsyncMock()
|
||||
session.exec.return_value.rowcount = 1
|
||||
session.refresh = AsyncMock()
|
||||
receive = AsyncMock(
|
||||
side_effect=[
|
||||
(1000, "sat", "https://mint.test"),
|
||||
ValueError("Mint Error: proofs already spent (Code: 11001)"),
|
||||
]
|
||||
)
|
||||
store = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.recieve_token", receive),
|
||||
patch("routstr.wallet.store_cashu_transaction", store),
|
||||
):
|
||||
results = await asyncio.gather(
|
||||
credit_balance("cashuAduplicate", key, session),
|
||||
credit_balance("cashuAduplicate", key, session),
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
assert sum(result == 1_000_000 for result in results) == 1
|
||||
failure = next(result for result in results if isinstance(result, Exception))
|
||||
classified = classify_redemption_error(failure)
|
||||
assert classified is not None and classified[3] == "cashu_token_already_spent"
|
||||
assert session.exec.await_count == 1
|
||||
store.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_credit_balance_constrains_redemption_to_key_mint() -> None:
|
||||
key_mint = "http://key-mint:3338"
|
||||
|
||||
Reference in New Issue
Block a user