Compare commits

...
Author SHA1 Message Date
9qeklajc 4fe8653537 add log 2026-07-11 23:13:50 +02:00
9qeklajc 3bd8bae543 Merge branch 'fix-too-low-request-costs' into token-storage-and-rate-limit 2026-07-11 22:31:31 +02:00
9qeklajc 527c4ae8a2 support new in/out token labels 2026-07-11 22:31:25 +02:00
9qeklajc 140af23c8e Merge branch 'fix-too-low-request-costs' into token-storage-and-rate-limit 2026-07-11 21:41:49 +02:00
9qeklajc d6de546279 fix too low prices calculation 2026-07-11 21:41:00 +02:00
9qeklajc 31ddfa96ad Merge branch 'fix-too-low-request-costs' into token-storage-and-rate-limit 2026-07-11 21:25:49 +02:00
9qeklajc 0178df4d25 fix too low prices calculation 2026-07-11 21:25:42 +02:00
9qeklajc 4231c62729 Merge branch 'fix/mint-rate-limit-and-fallback' into token-storage-and-rate-limit 2026-07-11 00:11:16 +02:00
9qeklajc acb630f6cf refactor: adapt mint throttling to 429 responses 2026-07-10 23:54:05 +02:00
9qeklajc 6ace3b48c1 Merge branch 'fix/mint-rate-limit-and-fallback' into token-storage-and-rate-limit 2026-07-10 23:46:23 +02:00
9qeklajc 1230d528de fix: avoid rate limiting balance proof checks 2026-07-10 23:46:09 +02:00
9qeklajc d2641da38f fix: linearize mint URL migration 2026-07-10 23:07:34 +02:00
9qeklajc c82d66da87 Merge branch 'enforce-out-token-creation' into token-storage-and-rate-limit 2026-07-10 22:57:22 +02:00
9qeklajc 77d14d928c Merge branch 'fix/mint-rate-limit-and-fallback' into token-storage-and-rate-limit 2026-07-10 22:57:18 +02:00
9qeklajc d23c90b939 fix: type wallet test fixture 2026-07-10 21:50:50 +02:00
9qeklajc d8db2a3051 fix: harden mint rate limiting and fallback 2026-07-10 21:46:56 +02:00
9qeklajc be33d2ee1b fix build 2026-07-10 21:19:59 +02:00
9qeklajc 0bbbf902cd Merge origin/main into fix/mint-rate-limit-and-fallback 2026-07-10 21:12:54 +02:00
9qeklajc 7ed18a9d02 fix: per-mint rate limiting, trusted-mint fallback, and retry factory fix 2026-07-10 20:43:07 +02:00
9qeklajc a3a05d2d5c enforce out token creation 2026-07-10 01:48:10 +02:00
20 changed files with 2132 additions and 204 deletions
@@ -0,0 +1,21 @@
"""add mint_url to lightning_invoices
Revision ID: add_mint_url_li
Revises: d7e8f9a0b1c2
Create Date: 2026-07-10 23:07:19.000000
"""
import sqlalchemy as sa
from alembic import op
revision = "add_mint_url_li"
down_revision = "d7e8f9a0b1c2"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("lightning_invoices", sa.Column("mint_url", sa.String(), nullable=True))
def downgrade() -> None:
op.drop_column("lightning_invoices", "mint_url")
@@ -0,0 +1,75 @@
"""unique (token, type) constraint on cashu_transactions
Revision ID: d7e8f9a0b1c2
Revises: c6d7e8f9a0b1
Create Date: 2026-07-10 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "d7e8f9a0b1c2"
down_revision = "c6d7e8f9a0b1"
branch_labels = None
depends_on = None
_CONSTRAINT_NAME = "uq_cashu_transactions_token_type"
def _dedup_rows(conn: sa.Connection) -> None:
"""Delete duplicate (token, type) rows, keeping the oldest per group."""
conn.execute(
sa.text(
"DELETE FROM cashu_transactions "
"WHERE id IN ("
" SELECT id FROM ("
" SELECT id, ROW_NUMBER() OVER ("
" PARTITION BY token, type "
" ORDER BY created_at ASC, id ASC"
" ) AS rn "
" FROM cashu_transactions"
" ) WHERE rn > 1"
")"
)
)
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
_dedup_rows(conn)
existing_indexes = {
idx["name"] for idx in inspector.get_indexes("cashu_transactions")
}
existing_constraints = {
uc["name"]
for uc in inspector.get_unique_constraints("cashu_transactions")
}
if _CONSTRAINT_NAME not in existing_indexes and _CONSTRAINT_NAME not in existing_constraints:
op.create_index(
_CONSTRAINT_NAME,
"cashu_transactions",
["token", "type"],
unique=True,
)
def downgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
existing_indexes = {
idx["name"] for idx in inspector.get_indexes("cashu_transactions")
}
existing_constraints = {
uc["name"]
for uc in inspector.get_unique_constraints("cashu_transactions")
}
if _CONSTRAINT_NAME in existing_indexes:
op.drop_index(_CONSTRAINT_NAME, table_name="cashu_transactions")
elif _CONSTRAINT_NAME in existing_constraints:
op.drop_constraint(_CONSTRAINT_NAME, "cashu_transactions")
+11 -14
View File
@@ -15,7 +15,7 @@ from .core.db import (
AsyncSession,
CashuTransaction,
get_session,
store_cashu_transaction,
store_cashu_transaction_with_retry,
)
from .core.logging import get_logger
from .core.settings import settings
@@ -457,19 +457,16 @@ async def refund_wallet_endpoint(
await _refund_cache_set(bearer_value, result)
if "token" in result:
try:
await store_cashu_transaction(
token=result["token"],
amount=remaining_balance,
unit=key.refund_currency or "sat",
mint_url=key.refund_mint_url,
typ="out",
collected=False,
source="apikey",
api_key_hashed_key=key.hashed_key,
)
except Exception:
pass # store_cashu_transaction already logs
await store_cashu_transaction_with_retry(
token=result["token"],
amount=remaining_balance,
unit=key.refund_currency or "sat",
mint_url=key.refund_mint_url,
typ="out",
collected=False,
source="apikey",
api_key_hashed_key=key.hashed_key,
)
logger.info(
"refund_wallet_endpoint: refund successful",
+315 -6
View File
@@ -1,3 +1,5 @@
import asyncio
import json
import os
import pathlib
import sqlite3
@@ -10,7 +12,7 @@ from alembic import command
from alembic.config import Config
from alembic.util.exc import CommandError
from sqlalchemy import UniqueConstraint, delete
from sqlalchemy.exc import OperationalError
from sqlalchemy.exc import IntegrityError, OperationalError
from sqlalchemy.ext.asyncio.engine import create_async_engine
from sqlalchemy.orm import aliased
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
@@ -26,6 +28,81 @@ DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
engine = create_async_engine(DATABASE_URL, echo=False) # echo=True for debugging SQL
# Durable JSONL outbox for cashu transactions when the DB is down.
_cashu_outbox_lock = asyncio.Lock()
def _cashu_outbox_path() -> pathlib.Path:
"""Resolve the per-process cashu outbox file path."""
override = os.environ.get("CASHU_OUTBOX_PATH")
if override:
return pathlib.Path(override)
filename = f"cashu_outbox.{os.getpid()}.jsonl"
if DATABASE_URL.startswith("sqlite"):
raw = DATABASE_URL.split("///", 1)[-1]
if not raw or raw == ":memory:":
return pathlib.Path("/tmp").resolve() / filename
db_file = pathlib.Path(raw)
return db_file.parent.resolve() / filename
return pathlib.Path(filename).resolve()
async def append_to_cashu_outbox(
*,
token: str,
amount: int,
unit: str,
mint_url: str | None,
typ: str,
request_id: str | None,
collected: bool,
created_at: int | None,
source: str,
api_key_hashed_key: str | None,
) -> None:
"""Append a transaction payload to the durable outbox. Never raises."""
entry = {
"outbox_id": uuid.uuid4().hex,
"queued_at": int(time.time()),
"token": token,
"amount": amount,
"unit": unit,
"mint_url": mint_url,
"type": typ,
"request_id": request_id,
"collected": collected,
"created_at": created_at,
"source": source,
"api_key_hashed_key": api_key_hashed_key,
}
line = json.dumps(entry, separators=(",", ":")) + "\n"
async with _cashu_outbox_lock:
try:
path = _cashu_outbox_path()
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "a", encoding="utf-8") as fh:
fh.write(line)
fh.flush()
os.fsync(fh.fileno())
logger.warning(
"Cashu transaction spooled to durable outbox",
extra={"outbox_path": str(path), "type": typ, "request_id": request_id},
)
except Exception as outbox_exc:
logger.critical(
"cashu outbox append failed; token may be unrecoverable",
extra={
"error": str(outbox_exc),
"type": typ,
"request_id": request_id,
"token": token,
},
)
class ApiKey(SQLModel, table=True): # type: ignore
__tablename__ = "api_keys"
@@ -225,6 +302,9 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
default=None, description="Associated API key hash for topup operations"
)
purpose: str = Field(description="create or topup")
mint_url: str | None = Field(
default=None, description="Mint URL where the quote was created (fallback tracking)"
)
created_at: int = Field(
default_factory=lambda: int(time.time()), description="Unix timestamp"
)
@@ -246,6 +326,11 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
class CashuTransaction(SQLModel, table=True): # type: ignore
__tablename__ = "cashu_transactions"
__table_args__ = (
UniqueConstraint(
"token", "type", name="uq_cashu_transactions_token_type"
),
)
id: str = Field(
primary_key=True,
@@ -287,9 +372,19 @@ async def store_cashu_transaction(
created_at: int | None = None,
source: str = "x-cashu",
api_key_hashed_key: str | None = None,
) -> None:
) -> bool:
"""Persist a cashu transaction; idempotent on (token, type). Returns True on success."""
try:
async with create_session() as session:
existing = await session.exec(
select(CashuTransaction).where(
CashuTransaction.token == token,
CashuTransaction.type == typ,
)
)
if existing.first() is not None:
return True
tx = CashuTransaction(
token=token,
amount=amount,
@@ -304,11 +399,225 @@ async def store_cashu_transaction(
)
session.add(tx)
await session.commit()
except Exception as e:
logger.warning(
f"Failed to store cashu transaction: {e} (type={typ})",
extra={"error": str(e), "type": typ},
return True
except IntegrityError:
# A concurrent insert of the same (token, type) won the race.
async with create_session() as session:
existing = await session.exec(
select(CashuTransaction).where(
CashuTransaction.token == token,
CashuTransaction.type == typ,
)
)
if existing.first() is not None:
return True
logger.error(
f"Integrity error storing cashu transaction: non-duplicate violation (type={typ})",
extra={
"error": "non-duplicate IntegrityError",
"type": typ,
"request_id": request_id,
"amount": amount,
"unit": unit,
"mint_url": mint_url,
},
)
return False
except Exception as e:
logger.error(
f"Failed to store cashu transaction: {e} (type={typ})",
extra={
"error": str(e),
"type": typ,
"request_id": request_id,
"amount": amount,
"unit": unit,
"mint_url": mint_url,
"token_preview": (token[:30] + "...") if len(token) > 30 else token,
},
)
return False
async def store_cashu_transaction_with_retry(
token: str,
amount: int,
unit: str,
mint_url: str | None = None,
typ: str = "out",
request_id: str | None = None,
collected: bool = False,
created_at: int | None = None,
source: str = "x-cashu",
api_key_hashed_key: str | None = None,
max_retries: int = 3,
) -> bool:
"""Retry ``store_cashu_transaction`` with backoff; spool to the outbox on exhaustion."""
for attempt in range(max_retries):
ok = await store_cashu_transaction(
token=token,
amount=amount,
unit=unit,
mint_url=mint_url,
typ=typ,
request_id=request_id,
collected=collected,
created_at=created_at,
source=source,
api_key_hashed_key=api_key_hashed_key,
)
if ok:
return True
if attempt < max_retries - 1:
backoff = 0.5 * (2 ** attempt)
logger.warning(
"Retrying store_cashu_transaction",
extra={
"attempt": attempt + 1,
"backoff_seconds": backoff,
"type": typ,
"request_id": request_id,
},
)
await asyncio.sleep(backoff)
logger.critical(
"Cashu transaction spooled to outbox after retries exhausted",
extra={
"type": typ,
"request_id": request_id,
"amount": amount,
"unit": unit,
"mint_url": mint_url,
"token": token,
},
)
await append_to_cashu_outbox(
token=token,
amount=amount,
unit=unit,
mint_url=mint_url,
typ=typ,
request_id=request_id,
collected=collected,
created_at=created_at or int(time.time()),
source=source,
api_key_hashed_key=api_key_hashed_key,
)
return False
async def replay_cashu_outbox(path: pathlib.Path | None = None) -> int:
"""Replay spooled outbox entries into the DB. Returns the count persisted."""
if path is None:
path = _cashu_outbox_path()
if not path.exists():
return 0
# Hold the lock for the full read-replay-rewrite cycle so a concurrent
# appender cannot slip a new line between the read and the rewrite.
async with _cashu_outbox_lock:
try:
raw_lines = path.read_text(encoding="utf-8").splitlines()
except Exception as e:
logger.error("Failed to read cashu outbox", extra={"error": str(e)})
return 0
entries: list[dict] = []
for line in raw_lines:
line = line.strip()
if not line:
continue
try:
entries.append(json.loads(line))
except json.JSONDecodeError as e:
logger.warning(
"Skipping malformed outbox line",
extra={"error": str(e), "line_preview": line[:80]},
)
if not entries:
return 0
persisted = 0
remaining: list[dict] = []
for entry in entries:
ok = await store_cashu_transaction(
token=entry["token"],
amount=entry["amount"],
unit=entry["unit"],
mint_url=entry.get("mint_url"),
typ=entry.get("type", "out"),
request_id=entry.get("request_id"),
collected=entry.get("collected", False),
created_at=entry.get("created_at"),
source=entry.get("source", "x-cashu"),
api_key_hashed_key=entry.get("api_key_hashed_key"),
)
if ok:
persisted += 1
logger.info(
"Outbox entry replayed into DB",
extra={
"outbox_id": entry.get("outbox_id"),
"type": entry.get("type"),
"request_id": entry.get("request_id"),
},
)
else:
remaining.append(entry)
try:
if remaining:
tmp = path.with_suffix(path.suffix + ".tmp")
with open(tmp, "w", encoding="utf-8") as fh:
for entry in remaining:
fh.write(json.dumps(entry, separators=(",", ":")) + "\n")
fh.flush()
os.fsync(fh.fileno())
os.replace(tmp, path)
else:
tmp = path.with_suffix(path.suffix + ".tmp")
tmp.write_text("")
os.replace(tmp, path)
except Exception as e:
logger.error(
"Failed to rewrite cashu outbox after replay",
extra={"error": str(e), "remaining": len(remaining)},
)
return persisted
async def replay_all_outbox_files() -> int:
"""Replay every cashu outbox file in the outbox directory."""
override = os.environ.get("CASHU_OUTBOX_PATH")
if override:
return await replay_cashu_outbox(pathlib.Path(override))
outbox_dir = _cashu_outbox_path().parent
total = 0
for file in sorted(outbox_dir.glob("cashu_outbox.*.jsonl")):
total += await replay_cashu_outbox(file)
return total
async def periodic_cashu_outbox_replay() -> None:
"""Background loop that drains the cashu outbox into the database."""
interval = float(os.environ.get("CASHU_OUTBOX_REPLAY_INTERVAL", "30"))
logger.info("Starting cashu outbox replay loop", extra={"interval": interval})
while True:
try:
await asyncio.sleep(interval)
await replay_all_outbox_files()
except asyncio.CancelledError:
raise
except Exception as e:
logger.error(
"Outbox replay iteration failed",
extra={"error": str(e), "error_type": type(e).__name__},
)
class UpstreamProviderRow(SQLModel, table=True): # type: ignore
+7 -1
View File
@@ -32,7 +32,7 @@ from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
from ..upstream.litellm_routing import configure_litellm
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
from .admin import admin_router
from .db import create_session, init_db, run_migrations
from .db import create_session, init_db, periodic_cashu_outbox_replay, run_migrations
from .exceptions import general_exception_handler, http_exception_handler
from .logging import get_logger, setup_logging
from .middleware import LoggingMiddleware
@@ -65,6 +65,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
refund_sweep_task = None
routstr_fee_task = None
invoice_watcher_task = None
cashu_outbox_task = None
try:
# Apply litellm-wide settings (drop_params, chat-completions URL,
@@ -142,6 +143,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher())
cashu_outbox_task = asyncio.create_task(periodic_cashu_outbox_replay())
yield
@@ -187,6 +189,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
routstr_fee_task.cancel()
if invoice_watcher_task is not None:
invoice_watcher_task.cancel()
if cashu_outbox_task is not None:
cashu_outbox_task.cancel()
try:
tasks_to_wait = []
@@ -220,6 +224,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(routstr_fee_task)
if invoice_watcher_task is not None:
tasks_to_wait.append(invoice_watcher_task)
if cashu_outbox_task is not None:
tasks_to_wait.append(cashu_outbox_task)
if tasks_to_wait:
await asyncio.gather(*tasks_to_wait, return_exceptions=True)
+22 -5
View File
@@ -49,6 +49,18 @@ class Settings(BaseSettings):
payout_interval_seconds: int = Field(
default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS"
)
# Timeout (seconds) for individual mint API operations (melt, mint, swap,
# checkstate). When a mint is slow or rate-limiting, operations are
# cancelled after this delay instead of hanging indefinitely.
mint_operation_timeout_seconds: int = Field(
default=30, gt=0, env="MINT_OPERATION_TIMEOUT_SECONDS"
)
# Maximum concurrent API operations per mint. Actual mint quotas vary by
# endpoint, so 429 responses drive adaptive cooldown instead of fixed RPM
# pacing. 0 = unlimited concurrency.
mint_max_concurrency: int = Field(default=4, ge=0, env="MINT_MAX_CONCURRENCY")
# Max retries when a mint returns 429 or times out (exponential backoff).
mint_retry_max_attempts: int = Field(default=3, ge=0, env="MINT_RETRY_MAX_ATTEMPTS")
# Pricing
# Default behavior: derive pricing from MODELS
@@ -97,7 +109,9 @@ class Settings(BaseSettings):
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(
default=604800, env="REFUND_SWEEP_TTL_SECONDS"
)
# Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL")
@@ -116,9 +130,8 @@ class Settings(BaseSettings):
# Discovery
relays: list[str] = Field(default_factory=list, env="RELAYS")
enable_analytics_sharing: bool = Field(
default=True, env="ENABLE_ANALYTICS_SHARING"
)
enable_analytics_sharing: bool = Field(default=True, env="ENABLE_ANALYTICS_SHARING")
def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
"""Discard unknown keys from persisted settings."""
@@ -281,7 +294,11 @@ class SettingsService:
valid_fields = set(env_resolved.dict().keys())
merged_dict: dict[str, Any] = dict(env_resolved.dict())
merged_dict.update(
{k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields}
{
k: v
for k, v in db_json.items()
if v not in (None, "", [], {}) and k in valid_fields
}
)
merged_dict = Settings(**merged_dict).dict()
+90 -18
View File
@@ -11,7 +11,13 @@ from sqlmodel.ext.asyncio.session import AsyncSession
from .core.db import ApiKey, LightningInvoice, create_session, get_session
from .core.logging import get_logger
from .core.settings import settings
from .wallet import get_wallet
from .wallet import (
MintConnectionError,
_is_mint_rate_limited,
_mint_operation,
get_wallet,
is_mint_connection_error,
)
logger = get_logger(__name__)
@@ -64,12 +70,50 @@ class InvoiceRecoverRequest(BaseModel):
bolt11: str = Field(description="BOLT11 invoice string")
async def _request_mint_with_fallback(
amount_sats: int,
*,
allowed_mints: list[str] | None = None,
) -> tuple[str, str, str]:
"""Request a quote, falling back only among the allowed trusted mints."""
tried: list[str] = []
configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints]
candidates = list(dict.fromkeys(configured))
for mint_url in candidates:
try:
wallet = await get_wallet(mint_url, "sat")
quote = await _mint_operation(
lambda: wallet.request_mint(amount_sats),
op_name="request_mint_invoice",
mint_url=mint_url,
)
return quote.request, quote.quote, mint_url
except Exception as e:
tried.append(f"{mint_url}: {type(e).__name__}")
if not is_mint_connection_error(e) and not _is_mint_rate_limited(e):
raise
logger.warning(
"request_mint failed, trying fallback mint",
extra={
"failed_mint": mint_url,
"error": str(e),
"tried": tried,
},
)
continue
raise MintConnectionError(f"All mints failed for request_mint: {tried}")
async def generate_lightning_invoice(
amount_sats: int, description: str
) -> tuple[str, str]:
wallet = await get_wallet(settings.primary_mint, "sat")
quote = await wallet.request_mint(amount_sats)
return quote.request, quote.quote
amount_sats: int,
description: str,
*,
allowed_mints: list[str] | None = None,
) -> tuple[str, str, str]:
bolt11, payment_hash, mint_url = await _request_mint_with_fallback(
amount_sats, allowed_mints=allowed_mints
)
return bolt11, payment_hash, mint_url
def generate_invoice_id() -> str:
@@ -83,6 +127,7 @@ async def create_invoice(
session: AsyncSession = Depends(get_session),
) -> InvoiceCreateResponse:
api_key_token = _extract_bearer_api_key(authorization) or request.api_key
topup_api_key: ApiKey | None = None
if request.purpose == "topup":
if not api_key_token:
@@ -93,14 +138,21 @@ async def create_invoice(
if not api_key_token.startswith("sk-"):
raise HTTPException(status_code=400, detail="Invalid API key format")
api_key = await session.get(ApiKey, api_key_token[3:])
if not api_key:
topup_api_key = await session.get(ApiKey, api_key_token[3:])
if not topup_api_key:
raise HTTPException(status_code=404, detail="API key not found")
try:
description = f"Routstr {request.purpose} {request.amount_sats} sats"
bolt11, payment_hash = await generate_lightning_invoice(
request.amount_sats, description
# An API key is backed by one mint. A top-up must use that same mint;
# falling back to another would create mixed-mint collateral that the
# current single refund_mint_url field cannot account for or refund.
allowed_mints = None
if request.purpose == "topup":
assert topup_api_key is not None
allowed_mints = [topup_api_key.refund_mint_url or settings.primary_mint]
bolt11, payment_hash, mint_url = await generate_lightning_invoice(
request.amount_sats, description, allowed_mints=allowed_mints
)
invoice_id = generate_invoice_id()
@@ -115,6 +167,7 @@ async def create_invoice(
status="pending",
api_key_hash=api_key_token[3:] if api_key_token else None,
purpose=request.purpose,
mint_url=mint_url,
balance_limit=request.balance_limit,
balance_limit_reset=request.balance_limit_reset,
validity_date=request.validity_date,
@@ -223,9 +276,14 @@ async def check_invoice_payment(
invoice: LightningInvoice, session: AsyncSession
) -> None:
try:
wallet = await get_wallet(settings.primary_mint, "sat")
mint_url = invoice.mint_url or settings.primary_mint
wallet = await get_wallet(mint_url, "sat")
mint_status = await wallet.get_mint_quote(invoice.payment_hash)
mint_status = await _mint_operation(
lambda: wallet.get_mint_quote(invoice.payment_hash),
op_name="get_mint_quote",
mint_url=mint_url,
)
if mint_status.paid:
invoice.status = "paid"
@@ -258,8 +316,14 @@ async def check_invoice_payment(
async def create_api_key_from_invoice(
invoice: LightningInvoice, session: AsyncSession
) -> ApiKey:
wallet = await get_wallet(settings.primary_mint, "sat")
await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash)
mint_url = invoice.mint_url or settings.primary_mint
wallet = await get_wallet(mint_url, "sat")
await _mint_operation(
lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash),
op_name="invoice_mint_create",
mint_url=mint_url,
retry_timeouts=False,
)
dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}"
hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest()
@@ -268,7 +332,7 @@ async def create_api_key_from_invoice(
hashed_key=hashed_key,
balance=invoice.amount_sats * 1000, # Convert to msats
refund_currency="sat",
refund_mint_url=settings.primary_mint,
refund_mint_url=mint_url,
balance_limit=invoice.balance_limit,
balance_limit_reset=invoice.balance_limit_reset,
validity_date=invoice.validity_date,
@@ -283,8 +347,14 @@ async def create_api_key_from_invoice(
async def topup_api_key_from_invoice(
invoice: LightningInvoice, session: AsyncSession
) -> None:
wallet = await get_wallet(settings.primary_mint, "sat")
await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash)
mint_url = invoice.mint_url or settings.primary_mint
wallet = await get_wallet(mint_url, "sat")
await _mint_operation(
lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash),
op_name="invoice_mint_topup",
mint_url=mint_url,
retry_timeouts=False,
)
if not invoice.api_key_hash:
raise ValueError("No API key associated with topup invoice")
@@ -297,7 +367,9 @@ async def topup_api_key_from_invoice(
await session.flush()
INVOICE_WATCH_INTERVAL_SECONDS = 5
# Nutshell mints throttle Lightning backend lookups to once per 10s per
# quote, so polling faster just burns the global request budget for nothing.
INVOICE_WATCH_INTERVAL_SECONDS = 10
INVOICE_WATCH_BATCH_LIMIT = 100
+125 -24
View File
@@ -1,4 +1,5 @@
import math
from decimal import ROUND_CEILING, ROUND_FLOOR, Decimal
from pydantic.v1 import BaseModel
@@ -157,11 +158,18 @@ async def calculate_cost(
},
)
try:
input_usd = _coerce_usd(
usage_data.get("cost_details", {}).get("input_cost", 0)
cost_details = usage_data.get("cost_details", {})
if not isinstance(cost_details, dict):
cost_details = {}
input_usd = _first_usd(
cost_details,
"input_cost",
"upstream_inference_prompt_cost",
)
output_usd = _coerce_usd(
usage_data.get("cost_details", {}).get("output_cost", 0)
output_usd = _first_usd(
cost_details,
"output_cost",
"upstream_inference_completions_cost",
)
return _calculate_from_usd_cost(
usd_cost,
@@ -256,6 +264,15 @@ def _coerce_usd(value: object) -> float:
return 0.0
def _first_usd(source: dict, *fields: str) -> float:
"""Return the first positive USD value among equivalent provider fields."""
for field in fields:
value = _coerce_usd(source.get(field))
if value > 0:
return value
return 0.0
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
"""Resolve USD cost with clear priority order.
@@ -263,7 +280,11 @@ def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
"""
cost_details = usage_data.get("cost_details")
if isinstance(cost_details, dict):
cost = _coerce_usd(cost_details.get("total_cost"))
cost = _first_usd(
cost_details,
"total_cost",
"upstream_inference_cost",
)
if cost > 0:
return cost
@@ -347,6 +368,43 @@ def _resolve_provider_fee(model_id: str) -> float:
return float(providers[0].provider_fee)
def _log_zero_cost_components(
cost: CostData,
response_data: dict,
calculation_source: str,
) -> None:
"""Log suspicious zero component costs without changing the billed total."""
input_usage = (
cost.input_tokens
+ cost.cache_read_input_tokens
+ cost.cache_creation_input_tokens
)
zero_components = []
if cost.total_msats > 0 and input_usage > 0 and cost.input_msats == 0:
zero_components.append("input")
if cost.total_msats > 0 and cost.output_tokens > 0 and cost.output_msats == 0:
zero_components.append("output")
if not zero_components:
return
logger.error(
"Positive token usage produced a zero millisatoshi cost component",
extra={
"zero_components": zero_components,
"calculation_source": calculation_source,
"model": response_data.get("model", "unknown"),
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"cache_read_input_tokens": cost.cache_read_input_tokens,
"cache_creation_input_tokens": cost.cache_creation_input_tokens,
"input_msats": cost.input_msats,
"output_msats": cost.output_msats,
"total_msats": cost.total_msats,
"total_usd": cost.total_usd,
},
)
def _calculate_from_usd_cost(
usd_cost: float,
input_usd: float,
@@ -359,27 +417,66 @@ def _calculate_from_usd_cost(
) -> CostData:
"""Calculate cost from USD figures, deriving input/output split from tokens."""
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
usd_cost = usd_cost * provider_fee
input_usd = input_usd * provider_fee
output_usd = output_usd * provider_fee
sats_per_usd = 1.0 / sats_usd_price()
cost_in_sats = usd_cost * sats_per_usd
cost_in_msats = math.ceil(cost_in_sats * 1000)
fee_decimal = Decimal(str(provider_fee))
usd_cost_decimal = Decimal(str(usd_cost)) * fee_decimal
input_usd_decimal = Decimal(str(input_usd)) * fee_decimal
output_usd_decimal = Decimal(str(output_usd)) * fee_decimal
sats_usd_decimal = Decimal(str(sats_usd_price()))
if input_usd > 0 or output_usd > 0:
input_msats = int((input_usd * sats_per_usd) * 1000)
output_msats = int((output_usd * sats_per_usd) * 1000)
else:
effective_input_tokens = (
input_tokens + cache_read_tokens + cache_creation_tokens
usd_cost = float(usd_cost_decimal)
input_usd = float(input_usd_decimal)
output_usd = float(output_usd_decimal)
cost_in_sats = float(usd_cost_decimal / sats_usd_decimal)
cost_in_msats = int(
(usd_cost_decimal * Decimal(1000) / sats_usd_decimal).to_integral_value(
rounding=ROUND_CEILING
)
total_tokens = effective_input_tokens + output_tokens
input_msats = (
int(cost_in_msats * effective_input_tokens / total_tokens)
if total_tokens > 0
else 0
)
if input_usd_decimal > 0 or output_usd_decimal > 0:
# The total is the authoritative billed amount. Allocating that integer
# total proportionally avoids losing sub-millisatoshi remainders when
# input and output components are each truncated independently.
component_usd = input_usd_decimal + output_usd_decimal
input_msats = int(
(
Decimal(cost_in_msats) * input_usd_decimal / component_usd
).to_integral_value(rounding=ROUND_FLOOR)
)
output_msats = cost_in_msats - input_msats
else:
# Providers often report only a total USD cost. Derive the visible
# input/output split from the model's relative token prices; raw token
# counts alone are misleading when completion tokens cost more.
try:
pricing_rates = _get_pricing_rates(response_data)
except ValueError:
pricing_rates = None
if pricing_rates is None:
input_rate = float(settings.fixed_per_1k_input_tokens) * 1000.0
output_rate = float(settings.fixed_per_1k_output_tokens) * 1000.0
cache_read_rate = input_rate
cache_creation_rate = input_rate
else:
input_rate, output_rate, cache_read_rate, cache_creation_rate = (
pricing_rates
)
input_weight = (
input_tokens * input_rate
+ cache_read_tokens * cache_read_rate
+ cache_creation_tokens * cache_creation_rate
)
output_weight = output_tokens * output_rate
total_weight = input_weight + output_weight
if total_weight > 0:
input_msats = math.floor(cost_in_msats * input_weight / total_weight)
output_msats = cost_in_msats - input_msats
else:
input_msats = 0
output_msats = cost_in_msats
logger.info(
"Using cost from usage data/details",
@@ -391,7 +488,7 @@ def _calculate_from_usd_cost(
},
)
return CostData(
cost = CostData(
base_msats=0,
input_msats=input_msats,
output_msats=output_msats,
@@ -404,6 +501,8 @@ def _calculate_from_usd_cost(
cache_read_msats=0,
cache_creation_msats=0,
)
_log_zero_cost_components(cost, response_data, "usd")
return cost
def _calculate_from_tokens(
@@ -460,7 +559,7 @@ def _calculate_from_tokens(
visible_output_msats = int(calc_output_msats)
visible_input_msats = token_based_cost - visible_output_msats
return CostData(
cost = CostData(
base_msats=0,
input_msats=visible_input_msats,
output_msats=visible_output_msats,
@@ -473,3 +572,5 @@ def _calculate_from_tokens(
cache_read_msats=int(calc_cache_read_msats),
cache_creation_msats=int(calc_cache_write_msats),
)
_log_zero_cost_components(cost, response_data, "tokens")
return cost
+20 -8
View File
@@ -9,7 +9,8 @@ from cashu.wallet.wallet import Proof, Wallet
# The Cashu library issues POST /v1/melt/bolt11 with timeout=None, so a hung or
# very slow mint can block a melt (and any caller, e.g. the payout loop)
# indefinitely. Bound it here so callers fail instead of hanging forever.
# indefinitely. _mint_operation (imported lazily in raw_send_to_lnurl to avoid
# a circular import with wallet.py) bounds it via MINT_OPERATION_TIMEOUT_SECONDS.
MELT_TIMEOUT_SECONDS = 60
try:
@@ -221,22 +222,33 @@ async def raw_send_to_lnurl(
lnurl_data["callback_url"], final_amount
)
melt_quote_resp = await wallet.melt_quote(invoice=bolt11_invoice)
from ..wallet import _mint_operation
melt_quote_resp = await _mint_operation(
lambda: wallet.melt_quote(invoice=bolt11_invoice),
op_name="lnurl_melt_quote",
mint_url=str(wallet.url),
)
if amount:
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
try:
_ = await asyncio.wait_for(
wallet.melt(
proofs=proofs,
invoice=bolt11_invoice,
fee_reserve_sat=melt_quote_resp.fee_reserve,
quote_id=melt_quote_resp.quote,
_mint_operation(
lambda: wallet.melt(
proofs=proofs,
invoice=bolt11_invoice,
fee_reserve_sat=melt_quote_resp.fee_reserve,
quote_id=melt_quote_resp.quote,
),
op_name="lnurl_melt",
mint_url=str(wallet.url),
retry_timeouts=False,
),
timeout=MELT_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError as e:
except (httpx.TimeoutException, asyncio.TimeoutError) as e:
raise LNURLError(
f"Melt timed out after {MELT_TIMEOUT_SECONDS}s (mint unresponsive)"
) from e
+60 -50
View File
@@ -20,7 +20,7 @@ from ..core.db import (
AsyncSession,
UpstreamProviderRow,
create_session,
store_cashu_transaction,
store_cashu_transaction_with_retry,
)
from ..core.exceptions import UpstreamError
from ..core.redaction import redact_org_ids
@@ -1748,14 +1748,19 @@ class BaseUpstreamProvider:
total_cost = max(
total_cost,
_coerce_usd(cd.get("total_cost")),
_coerce_usd(cd.get("upstream_inference_cost")),
)
input_cost = max(
input_cost,
_coerce_usd(cd.get("input_cost")),
_coerce_usd(cd.get("upstream_inference_prompt_cost")),
)
output_cost = max(
output_cost,
_coerce_usd(cd.get("output_cost")),
_coerce_usd(
cd.get("upstream_inference_completions_cost")
),
)
for field in ("total_cost", "cost"):
total_cost = max(
@@ -3286,7 +3291,7 @@ class BaseUpstreamProvider:
)
try:
await store_cashu_transaction(
await store_cashu_transaction_with_retry(
token=refund_token,
amount=amount,
unit=unit,
@@ -3294,8 +3299,25 @@ class BaseUpstreamProvider:
typ="out",
request_id=request_id,
)
except Exception:
pass # store_cashu_transaction already logs
except Exception as store_exc:
# store_cashu_transaction_with_retry returns False (and
# spools to the outbox) on normal DB failures, so this only
# fires on an unexpected raise. Catch it here so the outer
# retry loop does NOT re-mint a second token (which would
# double-spend). The full token is logged so it can be
# recovered manually if the outbox write also failed.
logger.critical(
"send_refund: store_cashu_transaction_with_retry raised — "
"refund token is minted; spooled to outbox if possible",
extra={
"error": str(store_exc),
"amount": amount,
"unit": unit,
"mint": mint,
"request_id": request_id,
"token": refund_token,
},
)
return refund_token
except Exception as e:
@@ -3640,17 +3662,14 @@ class BaseUpstreamProvider:
emergency_refund = amount
refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint)
response.headers["X-Cashu"] = refund_token
try:
await store_cashu_transaction(
token=refund_token,
amount=emergency_refund,
unit=unit,
mint_url=mint,
typ="out",
request_id=request_id,
)
except Exception:
pass
await store_cashu_transaction_with_retry(
token=refund_token,
amount=emergency_refund,
unit=unit,
mint_url=mint,
typ="out",
request_id=request_id,
)
logger.warning(
"Emergency refund issued due to JSON parse error",
@@ -4002,18 +4021,15 @@ class BaseUpstreamProvider:
headers = self.prepare_headers(dict(request.headers))
request_id = getattr(request.state, "request_id", None)
try:
await store_cashu_transaction(
token=x_cashu_token,
amount=amount,
unit=unit,
mint_url=mint,
typ="in",
request_id=request_id,
collected=True,
)
except Exception:
pass
await store_cashu_transaction_with_retry(
token=x_cashu_token,
amount=amount,
unit=unit,
mint_url=mint,
typ="in",
request_id=request_id,
collected=True,
)
logger.info(
"X-Cashu token redeemed for Responses API",
@@ -4604,17 +4620,14 @@ class BaseUpstreamProvider:
emergency_refund = amount
refund_token = await send_token(emergency_refund, unit=unit, mint_url=mint)
response.headers["X-Cashu"] = refund_token
try:
await store_cashu_transaction(
token=refund_token,
amount=emergency_refund,
unit=unit,
mint_url=mint,
typ="out",
request_id=request_id,
)
except Exception:
pass
await store_cashu_transaction_with_retry(
token=refund_token,
amount=emergency_refund,
unit=unit,
mint_url=mint,
typ="out",
request_id=request_id,
)
logger.warning(
"Emergency refund issued for Responses API due to JSON parse error",
@@ -4678,18 +4691,15 @@ class BaseUpstreamProvider:
headers = self.prepare_headers(dict(request.headers))
request_id = getattr(request.state, "request_id", None)
try:
await store_cashu_transaction(
token=x_cashu_token,
amount=amount,
unit=unit,
mint_url=mint,
typ="in",
request_id=request_id,
collected=True,
)
except Exception:
pass
await store_cashu_transaction_with_retry(
token=x_cashu_token,
amount=amount,
unit=unit,
mint_url=mint,
typ="in",
request_id=request_id,
collected=True,
)
logger.info(
"X-Cashu token redeemed successfully",
+28
View File
@@ -317,6 +317,27 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
total_cost += _coerce_float(usage.get("total_cost"))
input_cost += _coerce_float(usage.get("input_cost"))
output_cost += _coerce_float(usage.get("output_cost"))
cost_details = usage.get("cost_details")
if isinstance(cost_details, dict):
total_cost = max(
total_cost,
_coerce_float(cost_details.get("total_cost")),
_coerce_float(cost_details.get("upstream_inference_cost")),
)
input_cost = max(
input_cost,
_coerce_float(cost_details.get("input_cost")),
_coerce_float(
cost_details.get("upstream_inference_prompt_cost")
),
)
output_cost = max(
output_cost,
_coerce_float(cost_details.get("output_cost")),
_coerce_float(
cost_details.get("upstream_inference_completions_cost")
),
)
msg_for_meta = event.get("message")
if isinstance(msg_for_meta, dict):
@@ -340,14 +361,21 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
total_cost = max(
total_cost,
_coerce_float(root_cost_details.get("total_cost")),
_coerce_float(root_cost_details.get("upstream_inference_cost")),
)
input_cost = max(
input_cost,
_coerce_float(root_cost_details.get("input_cost")),
_coerce_float(
root_cost_details.get("upstream_inference_prompt_cost")
),
)
output_cost = max(
output_cost,
_coerce_float(root_cost_details.get("output_cost")),
_coerce_float(
root_cost_details.get("upstream_inference_completions_cost")
),
)
event_type = str(event.get("type") or "")
+320 -43
View File
@@ -3,10 +3,10 @@ import re
import socket
import time
import typing
from typing import TypedDict
from typing import Any, Awaitable, Callable, TypedDict
import httpx
from cashu.core.base import Proof, Token
from cashu.core.base import MintQuote, Proof, Token
from cashu.core.mint_info import MintInfo as _CashuMintInfo
from cashu.wallet.helpers import deserialize_token_from_string
from cashu.wallet.wallet import Wallet
@@ -14,7 +14,7 @@ from pydantic_core import PydanticUndefined
from sqlmodel import col, select, update
from .core import db, get_logger
from .core.db import store_cashu_transaction
from .core.db import store_cashu_transaction_with_retry
from .core.settings import settings
from .payment.lnurl import raw_send_to_lnurl
@@ -62,6 +62,155 @@ _TRANSPORT_EXC_TYPES: tuple[type[BaseException], ...] = (
)
class _MintRateGuard:
"""Bound concurrency and adapt to actual per-mint 429 responses."""
_guards: dict[str, "_MintRateGuard"] = {}
@classmethod
def get(cls, mint_url: str) -> "_MintRateGuard | None":
concurrency = settings.mint_max_concurrency
if concurrency <= 0:
return None
guard = cls._guards.get(mint_url)
if guard is None or guard._max_concurrency != concurrency:
guard = cls(mint_url, concurrency)
cls._guards[mint_url] = guard
return guard
def __init__(self, mint_url: str, max_concurrency: int):
self._mint_url = mint_url
self._max_concurrency = max_concurrency
self._semaphore = asyncio.Semaphore(max_concurrency)
self._cooldown_until = 0.0
def apply_cooldown(self, delay: float) -> None:
self._cooldown_until = max(
self._cooldown_until, time.monotonic() + max(0.0, delay)
)
async def run(self, factory: Callable[[], Awaitable[Any]]) -> Any:
async with self._semaphore:
wait = self._cooldown_until - time.monotonic()
if wait > 0:
logger.debug(
"Mint rate guard: cooling down",
extra={
"mint_url": self._mint_url,
"wait_seconds": round(wait, 2),
},
)
await asyncio.sleep(wait)
return await factory()
def _is_mint_rate_limited(error: BaseException) -> bool:
"""True if the mint returned a 429 or rate-limit indication."""
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, httpx.HTTPStatusError):
if current.response.status_code == 429:
return True
lowered = str(current).lower()
if "rate limit" in lowered or "too many requests" in lowered:
return True
current = current.__cause__ or current.__context__
return False
async def _mint_operation(
factory: Callable[[], Awaitable[Any]],
*,
op_name: str = "mint_operation",
mint_url: str = "",
retry_timeouts: bool = True,
) -> Any:
"""Run a mint operation with bounded concurrency and adaptive cooldown.
The timeout covers concurrency queueing, 429 cooldown, backoff, and network
work together. ``factory`` must return a fresh coroutine for every retry.
"""
guard = _MintRateGuard.get(mint_url) if mint_url else None
timeout = settings.mint_operation_timeout_seconds
max_attempts = settings.mint_retry_max_attempts + 1
async def invoke() -> Any:
if guard is not None:
return await guard.run(factory)
return await factory()
async def run_with_retries() -> Any:
for attempt in range(max_attempts):
try:
return await invoke()
except (asyncio.TimeoutError, httpx.TimeoutException) as exc:
if retry_timeouts and attempt < max_attempts - 1:
backoff = (2**attempt) + (time.monotonic() % 1.0)
logger.warning(
"Mint operation timed out, retrying",
extra={
"op_name": op_name,
"mint_url": mint_url,
"attempt": attempt + 1,
"backoff_seconds": round(backoff, 2),
},
)
await asyncio.sleep(backoff)
continue
raise httpx.TimeoutException(
f"{op_name} timed out (attempts: {attempt + 1})"
) from exc
except Exception as exc:
if not _is_mint_rate_limited(exc):
raise
backoff = (2**attempt) + (time.monotonic() % 1.0)
if isinstance(exc, httpx.HTTPStatusError):
retry_after = _parse_retry_after(exc.response.headers)
if retry_after is not None:
backoff = max(retry_after, backoff)
if guard is not None:
guard.apply_cooldown(backoff)
if attempt >= max_attempts - 1:
raise
logger.warning(
"Mint rate-limited, applying cooldown",
extra={
"op_name": op_name,
"mint_url": mint_url,
"attempt": attempt + 1,
"cooldown_seconds": round(backoff, 2),
},
)
if guard is None:
await asyncio.sleep(backoff)
raise RuntimeError(f"{op_name}: exhausted retries unexpectedly")
try:
if timeout > 0:
return await asyncio.wait_for(run_with_retries(), timeout=timeout)
return await run_with_retries()
except asyncio.TimeoutError as exc:
raise httpx.TimeoutException(
f"{op_name} exceeded its {timeout}s total timeout"
) from exc
def _parse_retry_after(headers: Any) -> float | None:
"""Parse a Retry-After header (delta-seconds form) into seconds."""
raw = headers.get("retry-after") or headers.get("Retry-After")
if raw is None:
return None
try:
return float(str(raw).strip())
except (TypeError, ValueError):
return None
def is_mint_connection_error(error: BaseException) -> bool:
"""True if ``error`` (or anything in its cause/context chain) is a mint
transport failure. Walks the chain because some sites re-raise transport
@@ -192,10 +341,19 @@ async def _redeem_same_mint(
that, not the face value, or routstr over-credits the user and its wallet
drifts insolvent.
"""
await wallet.load_mint(keyset_id=token_obj.keysets[0])
await _mint_operation(
lambda: wallet.load_mint(keyset_id=token_obj.keysets[0]),
op_name="redeem_load_mint",
mint_url=token_obj.mint,
)
wallet.verify_proofs_dleq(token_obj.proofs)
input_fees = wallet.get_fees_for_proofs(token_obj.proofs)
await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)
await _mint_operation(
lambda: wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True),
op_name="redeem_split",
mint_url=token_obj.mint,
retry_timeouts=False,
)
return int(token_obj.amount) - input_fees, token_obj.unit, token_obj.mint
@@ -237,7 +395,9 @@ async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int
all_mint_urls = list({k.mint_url for k in wallet.keysets.values()})
proof_summary = {
f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id)
f"{k.mint_url}/{k.unit.name}": sum(
p.amount for p in wallet.proofs if p.id == k.id
)
for k in wallet.keysets.values()
}
# Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet
@@ -341,6 +501,44 @@ def _melt_insufficient_shortfall(error: Exception) -> int | None:
return 1
async def _request_mint_with_fallback(
amount: int, *, op_name: str, primary_wallet: Wallet | None = None
) -> tuple[Wallet, str, MintQuote]:
"""Try request_mint on the primary mint, fall back to other trusted mints
on transport or rate-limit failure. Returns the wallet, mint_url, and quote."""
candidates = [settings.primary_mint] + [
m for m in settings.cashu_mints if m != settings.primary_mint
]
tried: list[str] = []
for mint_url in candidates:
try:
if mint_url == settings.primary_mint and primary_wallet is not None:
wallet = primary_wallet
else:
wallet = await get_wallet(mint_url, settings.primary_mint_unit)
quote = await _mint_operation(
lambda: wallet.request_mint(amount),
op_name=op_name,
mint_url=mint_url,
)
return wallet, mint_url, quote
except Exception as e:
tried.append(f"{mint_url}: {type(e).__name__}")
if not is_mint_connection_error(e) and not _is_mint_rate_limited(e):
raise
logger.warning(
"request_mint failed, trying fallback mint",
extra={
"failed_mint": mint_url,
"error": str(e),
"tried": tried,
"op_name": op_name,
},
)
continue
raise MintConnectionError(f"All mints failed for {op_name}: {tried}")
async def _calculate_swap_amount(
amount_msat: int,
token_unit: str,
@@ -374,8 +572,16 @@ async def _calculate_swap_amount(
)
try:
dummy_mint_quote = await primary_wallet.request_mint(receive_amount)
dummy_melt_quote = await token_wallet.melt_quote(dummy_mint_quote.request)
_, _, dummy_mint_quote = await _request_mint_with_fallback(
receive_amount,
op_name="swap_fee_est_mint_quote",
primary_wallet=primary_wallet,
)
dummy_melt_quote = await _mint_operation(
lambda: token_wallet.melt_quote(dummy_mint_quote.request),
op_name="swap_fee_est_melt_quote",
mint_url=token_mint_url,
)
fee_reserve = dummy_melt_quote.fee_reserve
input_fees = token_wallet.get_fees_for_proofs(proofs)
@@ -462,15 +668,27 @@ async def swap_to_primary_mint(
# amount recomputed from the fees the mint actually demands.
observed_extra_fee = 0
attempt = 0
dest_wallet = primary_wallet
dest_mint_url = settings.primary_mint
while True:
attempt += 1
mint_quote = await primary_wallet.request_mint(minted_amount)
dest_wallet, dest_mint_url, mint_quote = await _request_mint_with_fallback(
minted_amount, op_name="swap_request_mint", primary_wallet=primary_wallet
)
logger.info(
"swap_to_primary_mint: mint quote received",
extra={"mint_quote_id": mint_quote.quote, "attempt": attempt},
extra={
"mint_quote_id": mint_quote.quote,
"attempt": attempt,
"dest_mint": dest_mint_url,
},
)
melt_quote = await token_wallet.melt_quote(mint_quote.request)
melt_quote = await _mint_operation(
lambda: token_wallet.melt_quote(mint_quote.request),
op_name="swap_melt_quote",
mint_url=token_obj.mint,
)
input_fees = token_wallet.get_fees_for_proofs(token_obj.proofs)
total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees
logger.info(
@@ -523,11 +741,16 @@ async def swap_to_primary_mint(
continue
try:
_ = await token_wallet.melt(
proofs=token_obj.proofs,
invoice=mint_quote.request,
fee_reserve_sat=melt_quote.fee_reserve,
quote_id=melt_quote.quote,
_ = await _mint_operation(
lambda: token_wallet.melt(
proofs=token_obj.proofs,
invoice=mint_quote.request,
fee_reserve_sat=melt_quote.fee_reserve,
quote_id=melt_quote.quote,
),
op_name="swap_melt",
mint_url=token_obj.mint,
retry_timeouts=False,
)
except Exception as e:
# A down mint won't fix itself by retrying with a smaller amount.
@@ -576,14 +799,23 @@ async def swap_to_primary_mint(
break
logger.info(
"swap_to_primary_mint: melt succeeded, minting on primary",
extra={"minted_amount": minted_amount, "mint_quote_id": mint_quote.quote},
"swap_to_primary_mint: melt succeeded, minting on destination",
extra={
"minted_amount": minted_amount,
"mint_quote_id": mint_quote.quote,
"dest_mint": dest_mint_url,
},
)
await primary_wallet.load_proofs(reload=True)
pre_mint_balance = primary_wallet.available_balance.amount
await dest_wallet.load_proofs(reload=True)
pre_mint_balance = dest_wallet.available_balance.amount
try:
_ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote)
_ = await _mint_operation(
lambda: dest_wallet.mint(minted_amount, quote_id=mint_quote.quote),
op_name="swap_mint_on_primary",
mint_url=dest_mint_url,
retry_timeouts=False,
)
except Exception as e:
if "11003" in str(e) or "outputs already signed" in str(e).lower():
# Previous mint call signed outputs at the mint but failed before
@@ -591,13 +823,18 @@ async def swap_to_primary_mint(
# advance the counter so the next request derives fresh secrets.
logger.warning(
"swap_to_primary_mint: outputs already signed — recovering orphaned proofs",
extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount},
extra={
"mint_quote_id": mint_quote.quote,
"minted_amount": minted_amount,
},
)
try:
for keyset_id in primary_wallet.keysets:
await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25)
await primary_wallet.load_proofs(reload=True)
post_recovery_balance = primary_wallet.available_balance.amount
for keyset_id in dest_wallet.keysets:
await dest_wallet.restore_tokens_for_keyset(
keyset_id, to=1, batch=25
)
await dest_wallet.load_proofs(reload=True)
post_recovery_balance = dest_wallet.available_balance.amount
balance_gained = post_recovery_balance - pre_mint_balance
logger.info(
"swap_to_primary_mint: recovery scan completed",
@@ -648,14 +885,14 @@ async def swap_to_primary_mint(
"swap_to_primary_mint: completed successfully",
extra={
"foreign_mint": token_obj.mint,
"primary_mint": settings.primary_mint,
"dest_mint": dest_mint_url,
"original_amount": token_amount,
"minted_amount": minted_amount,
"unit": settings.primary_mint_unit,
},
)
return int(minted_amount), settings.primary_mint_unit, settings.primary_mint
return int(minted_amount), settings.primary_mint_unit, dest_mint_url
async def credit_balance(
@@ -734,7 +971,7 @@ async def credit_balance(
)
try:
await store_cashu_transaction(
await store_cashu_transaction_with_retry(
token=cashu_token,
amount=original_amount,
unit=original_unit,
@@ -744,7 +981,7 @@ async def credit_balance(
api_key_hashed_key=key.hashed_key,
)
except Exception:
pass
pass # store_cashu_transaction_with_retry already logs + spools to outbox
logger.debug(
"Cashu token successfully redeemed and stored",
@@ -760,18 +997,39 @@ async def credit_balance(
_wallets: dict[str, Wallet] = {}
_wallet_last_load: dict[str, float] = {}
_wallet_load_locks: dict[str, asyncio.Lock] = {}
# Minimum seconds between full mint info + proof reloads for the same
# wallet. Prevents redundant mint API calls when get_wallet(load=True)
# is called rapidly by multiple background tasks (balance fetch, payout,
# auto-topup all hitting get_wallet within the same cycle).
_WALLOAD_RELOAD_MIN_INTERVAL_SECONDS = 30
async def get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Wallet:
global _wallets
global _wallets, _wallet_last_load, _wallet_load_locks
id = f"{mint_url}_{unit}"
if id not in _wallets:
_wallets[id] = await Wallet.with_db(mint_url, db=".wallet", unit=unit)
lock = _wallet_load_locks.setdefault(id, asyncio.Lock())
async with lock:
if id not in _wallets:
_wallets[id] = await Wallet.with_db(mint_url, db=".wallet", unit=unit)
if load:
await _wallets[id].load_mint()
await _wallets[id].load_proofs(reload=True)
return _wallets[id]
if load:
now = time.monotonic()
last = _wallet_last_load.get(id, 0)
if now - last >= _WALLOAD_RELOAD_MIN_INTERVAL_SECONDS:
await _mint_operation(
lambda: _wallets[id].load_mint(),
op_name="load_mint",
mint_url=mint_url,
)
await _mint_operation(
lambda: _wallets[id].load_proofs(reload=True),
op_name="load_proofs",
mint_url=mint_url,
)
_wallet_last_load[id] = time.monotonic()
return _wallets[id]
def get_proofs_per_mint_and_unit(
@@ -793,15 +1051,28 @@ async def slow_filter_spend_proofs(proofs: list[Proof], wallet: Wallet) -> list[
return []
_proofs = []
_spent_proofs = []
for i in range(0, len(proofs), 1000):
pb = proofs[i : i + 1000]
proof_states = await wallet.check_proof_state(pb)
# Keep proof-state checks in large batches. Mint quotas count HTTP requests,
# so smaller batches make balance reads slower and more likely to hit 429s.
batch_size = 1000
for i in range(0, len(proofs), batch_size):
pb = proofs[i : i + batch_size]
proof_states = await _mint_operation(
lambda: wallet.check_proof_state(pb),
op_name="check_proof_state",
mint_url=str(wallet.url),
)
for proof, state in zip(pb, proof_states.states):
if str(state.state) != "spent":
_proofs.append(proof)
else:
_spent_proofs.append(proof)
await wallet.set_reserved_for_send(_spent_proofs, reserved=True)
if _spent_proofs:
await _mint_operation(
lambda: wallet.set_reserved_for_send(_spent_proofs, reserved=True),
op_name="set_reserved_spent_proofs",
mint_url=str(wallet.url),
retry_timeouts=False,
)
return _proofs
@@ -848,7 +1119,9 @@ async def fetch_all_balances(
"unit": unit,
"wallet_balance": proofs_balance,
"user_balance": user_balance,
"owner_balance": proofs_balance - user_balance if proofs_balance != 0 else 0,
"owner_balance": proofs_balance - user_balance
if proofs_balance != 0
else 0,
}
return result
except Exception as e:
@@ -1065,7 +1338,11 @@ async def periodic_routstr_fee_payout() -> None:
wallet, settings.primary_mint, "sat", not_reserved=True
)
amount_received = await raw_send_to_lnurl(
wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats
wallet,
proofs,
ROUTSTR_LN_ADDRESS,
"sat",
amount=accumulated_sats,
)
paid_msats = accumulated_sats * 1000
await db.reset_routstr_fee(session, paid_msats)
@@ -26,11 +26,17 @@ async def patch_invoice_generation() -> Any:
"""Stub out `generate_lightning_invoice` so no mint round-trip is needed."""
counter = {"n": 0}
async def fake_generate(amount_sats: int, description: str) -> tuple[str, str]:
async def fake_generate(
amount_sats: int,
description: str,
*,
allowed_mints: list[str] | None = None,
) -> tuple[str, str, str]:
counter["n"] += 1
return (
f"lnbc{amount_sats}n1pfakeinvoice{counter['n']}",
f"payment_hash_{counter['n']}",
"http://localhost:3338",
)
with patch(
@@ -95,6 +101,9 @@ async def test_topup_with_authorization_header(
body = resp.json()
assert body["amount_sats"] == 500
assert body["bolt11"].startswith("lnbc")
assert patch_invoice_generation.call_args.kwargs["allowed_mints"] == [
"http://localhost:3338"
]
@pytest.mark.integration
+10 -7
View File
@@ -235,7 +235,7 @@ async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> No
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store,
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)) as mock_store,
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
):
@@ -270,7 +270,7 @@ async def test_apikey_refund_logs_token() -> None:
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger") as mock_logger,
@@ -299,7 +299,7 @@ async def test_apikey_refund_log_includes_path() -> None:
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger") as mock_logger,
@@ -338,7 +338,7 @@ async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", mock_send_token),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
):
@@ -367,7 +367,10 @@ async def test_credit_balance_stores_apikey_transaction_history() -> None:
"routstr.wallet.recieve_token",
AsyncMock(return_value=(100, "sat", "https://mint.example")),
),
patch("routstr.wallet.store_cashu_transaction", AsyncMock()) as mock_store,
patch(
"routstr.wallet.store_cashu_transaction_with_retry",
AsyncMock(return_value=True),
) as mock_store,
):
amount = await credit_balance("cashuAtopup_token", key, session)
@@ -402,7 +405,7 @@ async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
"routstr.balance.send_token",
AsyncMock(side_effect=MintConnectionError("raw mint outage detail")),
),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger"),
@@ -437,7 +440,7 @@ async def test_apikey_refund_generic_failure_is_sanitized_500() -> None:
with (
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
patch("routstr.balance.send_token", AsyncMock(side_effect=RuntimeError(raw_error))),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
patch("routstr.balance.logger"),
+411
View File
@@ -0,0 +1,411 @@
"""Unit tests for the Cashu outbox + retry persistence path.
These cover the guarantees introduced by the refund-token persistence fix:
* ``store_cashu_transaction`` is idempotent (no duplicate rows on retry)
* ``store_cashu_transaction_with_retry`` retries, then spools the full
token to a durable outbox when the DB stays unavailable
* ``replay_cashu_outbox`` drains the outbox into the DB and is idempotent
"""
import asyncio
import json
import os
from collections.abc import AsyncIterator
from pathlib import Path
from typing import Any
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel, select
import routstr.core.db as db
from routstr.core.db import (
CashuTransaction,
create_session,
replay_all_outbox_files,
replay_cashu_outbox,
store_cashu_transaction,
store_cashu_transaction_with_retry,
)
def _in_memory_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def fresh_engine(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> AsyncIterator[tuple[AsyncEngine, Path]]:
"""Point db.engine at an isolated in-memory DB and the outbox at a tmp file."""
engine = _in_memory_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
monkeypatch.setattr(db, "engine", engine)
# create_session() looks up module-global `engine` at call time.
monkeypatch.setattr(db, "create_session", create_session, raising=False)
outbox = tmp_path / "cashu_outbox.jsonl"
monkeypatch.setenv("CASHU_OUTBOX_PATH", str(outbox))
try:
yield engine, outbox
finally:
await engine.dispose()
@pytest.mark.asyncio
async def test_store_is_idempotent_on_duplicate(
fresh_engine: tuple[AsyncEngine, Path]
) -> None:
"""Inserting the same (token, type) twice yields exactly one row."""
_, _ = fresh_engine
ok1 = await store_cashu_transaction(
token="cashuAtokenA", amount=100, unit="sat", typ="out", request_id="req-1"
)
ok2 = await store_cashu_transaction(
token="cashuAtokenA", amount=100, unit="sat", typ="out", request_id="req-1"
)
assert ok1 is True
assert ok2 is True # duplicate is treated as success, not a new row
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
assert rows[0].token == "cashuAtokenA"
@pytest.mark.asyncio
async def test_with_retry_succeeds_first_try(
fresh_engine: tuple[AsyncEngine, Path]
) -> None:
_, _ = fresh_engine
ok = await store_cashu_transaction_with_retry(
token="cashuAtokenB", amount=50, unit="sat", typ="out", request_id="req-2"
)
assert ok is True
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
@pytest.mark.asyncio
async def test_with_retry_succeeds_after_transient_failure(
fresh_engine: tuple[AsyncEngine, Path], monkeypatch: pytest.MonkeyPatch
) -> None:
"""First attempt fails, second succeeds; no outbox spool, one row."""
_, _ = fresh_engine
calls = {"n": 0}
real_store = store_cashu_transaction
async def flaky_store(*args: Any, **kwargs: Any) -> bool:
calls["n"] += 1
if calls["n"] == 1:
return False
return await real_store(*args, **kwargs)
monkeypatch.setattr(db, "store_cashu_transaction", flaky_store)
# zero backoff so the test is fast
monkeypatch.setattr(db.asyncio, "sleep", lambda *_a, **_k: _asyncio_sleep_zero())
ok = await store_cashu_transaction_with_retry(
token="cashuAtokenC", amount=50, unit="sat", typ="out", request_id="req-3",
max_retries=3,
)
assert ok is True
assert calls["n"] == 2
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
@pytest.mark.asyncio
async def test_with_retry_spools_full_token_to_outbox_on_exhaustion(
fresh_engine: tuple[AsyncEngine, Path], monkeypatch: pytest.MonkeyPatch
) -> None:
"""When the DB stays down, the full token is spooled to the outbox file."""
_, outbox = fresh_engine
async def always_fail(*args: Any, **kwargs: Any) -> bool:
return False
monkeypatch.setattr(db, "store_cashu_transaction", always_fail)
monkeypatch.setattr(db.asyncio, "sleep", lambda *_a, **_k: _asyncio_sleep_zero())
full_token = "cashuAtokenD_super_secret_full_value_xyz"
ok = await store_cashu_transaction_with_retry(
token=full_token, amount=77, unit="sat", typ="out", request_id="req-4",
max_retries=3,
)
assert ok is False
assert outbox.exists()
lines = outbox.read_text(encoding="utf-8").strip().splitlines()
assert len(lines) == 1
entry = json.loads(lines[0])
# The full token must be recoverable from the outbox, not a preview.
assert entry["token"] == full_token
assert entry["amount"] == 77
assert entry["request_id"] == "req-4"
assert entry["type"] == "out"
@pytest.mark.asyncio
async def test_replay_drains_outbox_into_db(
fresh_engine: tuple[AsyncEngine, Path]
) -> None:
"""Outbox entries are persisted on replay and the outbox is cleared."""
_, outbox = fresh_engine
entry = {
"outbox_id": "abc123",
"queued_at": 1,
"token": "cashuAtokenE",
"amount": 200,
"unit": "sat",
"mint_url": None,
"type": "out",
"request_id": "req-5",
"collected": False,
"created_at": None,
"source": "x-cashu",
"api_key_hashed_key": None,
}
outbox.write_text(json.dumps(entry) + "\n", encoding="utf-8")
persisted = await replay_cashu_outbox()
assert persisted == 1
# outbox should now be empty (truncated)
assert outbox.read_text(encoding="utf-8") == ""
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
assert rows[0].token == "cashuAtokenE"
@pytest.mark.asyncio
async def test_replay_is_idempotent(fresh_engine: tuple[AsyncEngine, Path]) -> None:
"""Replaying an outbox whose entries already exist in the DB drops them."""
_, outbox = fresh_engine
# Pre-seed the DB row.
await store_cashu_transaction(
token="cashuAtokenF", amount=10, unit="sat", typ="out", request_id="req-6"
)
# Outbox references the same token.
entry = {
"outbox_id": "xyz",
"queued_at": 1,
"token": "cashuAtokenF",
"amount": 10,
"unit": "sat",
"mint_url": None,
"type": "out",
"request_id": "req-6",
"collected": False,
"created_at": None,
"source": "x-cashu",
"api_key_hashed_key": None,
}
outbox.write_text(json.dumps(entry) + "\n", encoding="utf-8")
persisted = await replay_cashu_outbox()
assert persisted == 1 # counted as persisted (already-present = success)
assert outbox.read_text(encoding="utf-8") == ""
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1 # no duplicate
@pytest.mark.asyncio
async def test_replay_keeps_entries_that_still_fail(
fresh_engine: tuple[AsyncEngine, Path], monkeypatch: pytest.MonkeyPatch
) -> None:
"""If replay still can't persist an entry, it stays in the outbox."""
_, outbox = fresh_engine
monkeypatch.setattr(db, "store_cashu_transaction", lambda *a, **k: _afail())
monkeypatch.setattr(db.asyncio, "sleep", lambda *_a, **_k: _asyncio_sleep_zero())
entry = {
"outbox_id": "keepme",
"queued_at": 1,
"token": "cashuAtokenG",
"amount": 5,
"unit": "sat",
"mint_url": None,
"type": "out",
"request_id": "req-7",
"collected": False,
"created_at": None,
"source": "x-cashu",
"api_key_hashed_key": None,
}
outbox.write_text(json.dumps(entry) + "\n", encoding="utf-8")
persisted = await replay_cashu_outbox()
assert persisted == 0
remaining = outbox.read_text(encoding="utf-8").strip().splitlines()
assert len(remaining) == 1
assert json.loads(remaining[0])["outbox_id"] == "keepme"
@pytest.mark.asyncio
async def test_replay_drops_corrupt_lines_and_persists_valid_ones(
fresh_engine: tuple[AsyncEngine, Path]
) -> None:
"""Corrupt JSONL lines are dropped; valid entries are persisted."""
_, outbox = fresh_engine
valid_entry = {
"outbox_id": "good1",
"queued_at": 1,
"token": "cashuAtokenH",
"amount": 300,
"unit": "sat",
"mint_url": None,
"type": "out",
"request_id": "req-8",
"collected": False,
"created_at": None,
"source": "x-cashu",
"api_key_hashed_key": None,
}
outbox.write_text(
json.dumps(valid_entry) + "\n" + '{"token": "broken\n',
encoding="utf-8",
)
persisted = await replay_cashu_outbox()
assert persisted == 1
assert outbox.read_text(encoding="utf-8") == ""
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
assert rows[0].token == "cashuAtokenH"
@pytest.mark.asyncio
async def test_store_concurrent_duplicate_insert_yields_one_row(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
"""Two concurrent stores of the same (token, type) yield one row, both True."""
db_file = tmp_path / "concurrent.db"
engine = create_async_engine(f"sqlite+aiosqlite:///{db_file}")
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
monkeypatch.setattr(db, "engine", engine)
monkeypatch.setattr(db, "create_session", create_session, raising=False)
monkeypatch.setenv("CASHU_OUTBOX_PATH", str(tmp_path / "outbox.jsonl"))
try:
results = await asyncio.gather(
store_cashu_transaction(
token="cashuAtokenI", amount=100, unit="sat", typ="out", request_id="req-9"
),
store_cashu_transaction(
token="cashuAtokenI", amount=100, unit="sat", typ="out", request_id="req-9"
),
)
assert all(results)
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
assert rows[0].token == "cashuAtokenI"
finally:
await engine.dispose()
@pytest.mark.asyncio
async def test_replay_under_concurrent_replay_is_idempotent(
fresh_engine: tuple[AsyncEngine, Path]
) -> None:
"""Two concurrent replay calls on a seeded outbox yield one row, outbox empty."""
_, outbox = fresh_engine
entry = {
"outbox_id": "concurrent1",
"queued_at": 1,
"token": "cashuAtokenJ",
"amount": 42,
"unit": "sat",
"mint_url": None,
"type": "out",
"request_id": "req-10",
"collected": False,
"created_at": None,
"source": "x-cashu",
"api_key_hashed_key": None,
}
outbox.write_text(json.dumps(entry) + "\n", encoding="utf-8")
results = await asyncio.gather(replay_cashu_outbox(), replay_cashu_outbox())
# At least one call persisted the entry; both may report success.
assert sum(results) >= 1
assert outbox.read_text(encoding="utf-8") == ""
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
assert len(rows) == 1
assert rows[0].token == "cashuAtokenJ"
@pytest.mark.asyncio
async def test_replay_all_drains_orphaned_outbox_files(
monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
"""replay_all_outbox_files drains every per-PID outbox file, including orphans."""
engine = _in_memory_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
monkeypatch.setattr(db, "engine", engine)
monkeypatch.setattr(db, "create_session", create_session, raising=False)
# No CASHU_OUTBOX_PATH override; _cashu_outbox_path resolves into tmp_path.
monkeypatch.delenv("CASHU_OUTBOX_PATH", raising=False)
monkeypatch.setattr(
db, "_cashu_outbox_path",
lambda: tmp_path / f"cashu_outbox.{os.getpid()}.jsonl",
)
def _make_entry(token: str, oid: str) -> dict[str, object]:
return {
"outbox_id": oid,
"queued_at": 1,
"token": token,
"amount": 100,
"unit": "sat",
"mint_url": None,
"type": "out",
"request_id": f"req-{oid}",
"collected": False,
"created_at": None,
"source": "x-cashu",
"api_key_hashed_key": None,
}
orphan_a = tmp_path / "cashu_outbox.11111.jsonl"
orphan_b = tmp_path / "cashu_outbox.22222.jsonl"
orphan_a.write_text(json.dumps(_make_entry("cashuOrphanA", "a1")) + "\n", encoding="utf-8")
orphan_b.write_text(json.dumps(_make_entry("cashuOrphanB", "b2")) + "\n", encoding="utf-8")
try:
total = await replay_all_outbox_files()
assert total == 2
assert orphan_a.read_text(encoding="utf-8") == ""
assert orphan_b.read_text(encoding="utf-8") == ""
async with create_session() as session:
rows = (await session.exec(select(CashuTransaction))).all()
tokens = {r.token for r in rows}
assert tokens == {"cashuOrphanA", "cashuOrphanB"}
finally:
await engine.dispose()
# --- tiny async helpers (can't use `await` inside lambdas) ---
async def _asyncio_sleep_zero() -> None:
"""No-op async sleep replacement for fast tests."""
return None
async def _afail() -> bool:
return False
+125 -1
View File
@@ -5,7 +5,7 @@ edge cases, and billing accuracy.
"""
import os
from unittest.mock import patch
from unittest.mock import Mock, patch
import pytest
@@ -465,6 +465,130 @@ async def test_cache_read_only_usd_cost_response_is_billed(
assert result.cache_read_input_tokens == 1000
@pytest.mark.asyncio
@pytest.mark.parametrize(
("total_cost", "input_cost", "output_cost", "expected_msats"),
[
(0.000471, 0.00023451, 0.00023649, 9420),
(0.00000004, 0.00000002, 0.00000002, 1),
],
)
async def test_small_usd_cost_components_sum_to_rounded_total(
total_cost: float,
input_cost: float,
output_cost: float,
expected_msats: int,
) -> None:
"""Small USD component costs must retain every billed millisatoshi."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 1,
"completion_tokens": 1,
"cost_details": {
"total_cost": total_cost,
"input_cost": input_cost,
"output_cost": output_cost,
},
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.total_msats == expected_msats
assert result.input_msats + result.output_msats == result.total_msats
@pytest.mark.asyncio
async def test_zero_msat_component_with_positive_tokens_logs_error() -> None:
"""A rounded-to-zero component is observable even when billing is valid."""
response = {
"model": "gpt-4",
"usage": {
"prompt_tokens": 1,
"completion_tokens": 1,
"cost_details": {
"total_cost": 0.00000004,
"input_cost": 0.00000002,
"output_cost": 0.00000002,
},
},
}
with patch("routstr.payment.cost_calculation.logger.error") as log_error:
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.total_msats == 1
log_error.assert_called_once()
assert log_error.call_args.args[0] == (
"Positive token usage produced a zero millisatoshi cost component"
)
assert log_error.call_args.kwargs["extra"]["zero_components"] == ["input"]
@pytest.mark.asyncio
async def test_total_only_usd_cost_uses_model_prices_for_component_split(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A reported total is split by priced tokens, not raw token counts."""
monkeypatch.setattr(settings, "fixed_pricing", False)
response = {
"model": "z-ai/glm-5.2-20260616",
"usage": {
"prompt_tokens": 375,
"completion_tokens": 218,
"total_cost": 0.00039088,
},
}
pricing = Mock(
prompt=0.0001,
completion=0.001,
input_cache_read=0.0001,
input_cache_write=0.0001,
)
model = Mock(sats_pricing=pricing)
with patch("routstr.proxy.get_model_instance", return_value=model):
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.total_msats == 7818
assert result.input_msats + result.output_msats == result.total_msats
assert result.input_msats == 1147
assert result.output_msats == 6671
@pytest.mark.asyncio
async def test_upstream_inference_cost_details_set_nonzero_components() -> None:
"""OpenAI-compatible upstream inference aliases retain their exact split."""
response = {
"model": "z-ai/glm-5.2-20260616",
"usage": {
"prompt_tokens": 211,
"completion_tokens": 500,
"total_tokens": 711,
"cost": 0.00242155,
"cost_details": {
"upstream_inference_cost": 0.00242155,
"upstream_inference_prompt_cost": 0.00022155,
"upstream_inference_completions_cost": 0.0022,
},
},
}
result = await calculate_cost(response, max_cost=100000)
assert isinstance(result, CostData)
assert result.input_tokens == 211
assert result.output_tokens == 500
assert result.input_msats == 4431
assert result.output_msats == 44000
assert result.total_msats == 48431
assert result.input_msats + result.output_msats == result.total_msats
# ============================================================================
# Test 13: Missing Usage Block
# ============================================================================
@@ -161,6 +161,31 @@ def test_cost_details_extracted_from_event_root() -> None:
assert result.output_cost == 0.005
@pytest.mark.unit
def test_upstream_inference_cost_details_extracted_from_usage() -> None:
"""Provider inference aliases inside usage preserve the USD split."""
event = {
"usage": {
"input_tokens": 211,
"output_tokens": 500,
"cost": 0.00242155,
"cost_details": {
"upstream_inference_cost": 0.00242155,
"upstream_inference_prompt_cost": 0.00022155,
"upstream_inference_completions_cost": 0.0022,
},
}
}
result = annotate_event(event, None)
assert result.input_tokens == 211
assert result.output_tokens == 500
assert result.total_cost == 0.00242155
assert result.input_cost == 0.00022155
assert result.output_cost == 0.0022
# ============================================================================
# Test 7: No Duplicated Dict Lookups
# ============================================================================
+2 -2
View File
@@ -1481,7 +1481,7 @@ async def test_x_cashu_transport_error_after_redemption_is_not_retryable(
"routstr.upstream.base.recieve_token",
new=AsyncMock(return_value=(5_000, "sat", "https://mint")),
),
patch("routstr.upstream.base.store_cashu_transaction", new=AsyncMock()),
patch("routstr.upstream.base.store_cashu_transaction_with_retry", new=AsyncMock(return_value=True)),
patch.object(
provider,
forward_attr,
@@ -1601,7 +1601,7 @@ async def test_x_cashu_zero_value_rejected_not_forwarded(
"routstr.upstream.base.recieve_token",
new=AsyncMock(return_value=(amount, "sat", "https://mint")),
),
patch("routstr.upstream.base.store_cashu_transaction", new=AsyncMock()),
patch("routstr.upstream.base.store_cashu_transaction_with_retry", new=AsyncMock(return_value=True)),
patch.object(
provider,
forward_attr,
+1 -1
View File
@@ -215,7 +215,7 @@ async def test_reset_all_reserved_balances_clears_reserved_at(session: AsyncSess
def _refund_patches(refund_token: str = "cashuArefund"): # type: ignore[no-untyped-def]
return (
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
patch("routstr.balance.store_cashu_transaction_with_retry", AsyncMock(return_value=True)),
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance._refund_cache_set", AsyncMock()),
)
+454 -23
View File
@@ -1,6 +1,8 @@
import asyncio
import base64
import json
import socket
from collections.abc import Generator
from unittest.mock import AsyncMock, Mock, patch
import httpx
@@ -19,6 +21,26 @@ from routstr.wallet import (
)
@pytest.fixture(autouse=True)
def isolate_wallet_runtime_state() -> Generator[None, None, None]:
"""Keep production limiter/wallet caches from leaking across unit tests."""
from routstr import wallet as wallet_module
from routstr.core.settings import settings
original_concurrency = settings.mint_max_concurrency
settings.mint_max_concurrency = 0
wallet_module._MintRateGuard._guards.clear()
wallet_module._wallets.clear()
wallet_module._wallet_last_load.clear()
wallet_module._wallet_load_locks.clear()
yield
settings.mint_max_concurrency = original_concurrency
wallet_module._MintRateGuard._guards.clear()
wallet_module._wallets.clear()
wallet_module._wallet_last_load.clear()
wallet_module._wallet_load_locks.clear()
@pytest.mark.asyncio
async def test_get_balance() -> None:
mock_wallet = Mock()
@@ -728,9 +750,7 @@ async def test_calculate_swap_amount_same_mint_short_circuit() -> None:
quotes are requested."""
from routstr.wallet import _calculate_swap_amount
_, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(
1000, fee_reserves=[]
)
_, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(1000, fee_reserves=[])
from routstr.core.settings import settings
@@ -755,9 +775,7 @@ async def test_calculate_swap_amount_msat_primary_unit() -> None:
"""With an msat primary mint the dummy quote and result stay in msats."""
from routstr.wallet import _calculate_swap_amount
_, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(
179, fee_reserves=[2]
)
_, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(179, fee_reserves=[2])
from routstr.core.settings import settings
@@ -805,12 +823,8 @@ async def test_calculate_swap_amount_wraps_estimation_failure() -> None:
"""Estimation infrastructure failures surface as a single clear ValueError."""
from routstr.wallet import _calculate_swap_amount
_, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(
179, fee_reserves=[]
)
mock_primary_wallet.request_mint = AsyncMock(
side_effect=Exception("mint offline")
)
_, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(179, fee_reserves=[])
mock_primary_wallet.request_mint = AsyncMock(side_effect=Exception("mint offline"))
from routstr.core.settings import settings
@@ -1062,9 +1076,9 @@ async def test_credit_balance_msat_unit_not_converted() -> None:
@pytest.mark.asyncio
async def test_credit_balance_survives_audit_store_failure() -> None:
"""A failure writing the CashuTransaction history record must not undo the
already-committed balance credit. (The silent swallow is a known
audit-trail gap slated for its own fix — this test pins the financial
invariant that the user keeps their credit, not the swallow itself.)"""
already-committed balance credit. (Persistence is spooled to the outbox on
failure; this test pins the financial invariant that the user keeps their
credit, not the spool behavior itself.)"""
mock_key = Mock()
mock_key.balance = 0
mock_key.hashed_key = "test_hash"
@@ -1078,8 +1092,8 @@ async def test_credit_balance_survives_audit_store_failure() -> None:
return_value=(1000, "sat", "http://mint:3338"),
):
with patch(
"routstr.wallet.store_cashu_transaction",
side_effect=Exception("history table locked"),
"routstr.wallet.store_cashu_transaction_with_retry",
AsyncMock(return_value=False),
):
amount = await credit_balance("cashuAtest", mock_key, mock_session)
@@ -1158,7 +1172,9 @@ def test_is_mint_connection_error_detects_transport_failures(
ValueError("Invalid Cashu token"),
# Mint answered with an error status — reachable, so NOT a connection error.
httpx.HTTPStatusError(
"500", request=httpx.Request("POST", "http://m"), response=httpx.Response(500)
"500",
request=httpx.Request("POST", "http://m"),
response=httpx.Response(500),
),
RuntimeError("some internal fault"),
],
@@ -1213,7 +1229,9 @@ def test_classify_zero_value(error: ValueError) -> None:
def test_classify_generic_valueerror_is_not_zero_value() -> None:
"""A generic wallet ValueError still falls to the generic bucket — the
zero-value match must not over-trigger."""
classified = classify_redemption_error(ValueError("some unexpected wallet condition"))
classified = classify_redemption_error(
ValueError("some unexpected wallet condition")
)
assert classified is not None
type_, status, _msg, code = classified
assert (type_, status, code) == (
@@ -1276,7 +1294,9 @@ async def test_credit_balance_db_transport_error_is_token_consumed() -> None:
@pytest.mark.asyncio
async def test_swap_fee_estimation_transport_error_raises_mint_connection_error() -> None:
async def test_swap_fee_estimation_transport_error_raises_mint_connection_error() -> (
None
):
"""A transport failure while estimating fees is surfaced as
MintConnectionError (→ 503), not a generic fee ValueError (→ 422)."""
from routstr.wallet import swap_to_primary_mint
@@ -1308,9 +1328,7 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None:
mock_token, mock_token_wallet, mock_primary_wallet = _make_swap_mocks(
1000, fee_reserves=[10, 10]
)
mock_token_wallet.melt = AsyncMock(
side_effect=httpx.ConnectTimeout("timed out")
)
mock_token_wallet.melt = AsyncMock(side_effect=httpx.ConnectTimeout("timed out"))
from routstr.core.settings import settings
@@ -1321,3 +1339,416 @@ async def test_swap_melt_transport_error_raises_mint_connection_error() -> None:
await swap_to_primary_mint(mock_token, mock_token_wallet)
assert mock_token_wallet.melt.call_count == 1
# ---------------------------------------------------------------------------
# Per-mint adaptive guard + _mint_operation factory/retry
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_balance_proof_check_uses_large_batches_to_avoid_rate_limit() -> None:
"""Balance reads must not turn a few hundred proofs into many mint requests."""
from routstr.wallet import slow_filter_spend_proofs
proofs = [Mock() for _ in range(250)]
states = [Mock(state="UNSPENT") for _ in proofs]
wallet = Mock()
wallet.url = "http://mint:3338"
wallet.check_proof_state = AsyncMock(return_value=Mock(states=states))
wallet.set_reserved_for_send = AsyncMock()
result = await slow_filter_spend_proofs(proofs, wallet)
assert result == proofs
wallet.check_proof_state.assert_awaited_once_with(proofs)
wallet.set_reserved_for_send.assert_not_awaited()
@pytest.mark.asyncio
async def test_mint_rate_guard_bounds_concurrency() -> None:
from routstr.wallet import _MintRateGuard
guard = _MintRateGuard("http://mint:3338", 2)
active = 0
peak = 0
async def operation() -> None:
nonlocal active, peak
active += 1
peak = max(peak, active)
await asyncio.sleep(0)
active -= 1
await asyncio.gather(*(guard.run(operation) for _ in range(5)))
assert peak == 2
@pytest.mark.asyncio
async def test_mint_rate_guard_waits_for_adaptive_cooldown() -> None:
from routstr.wallet import _MintRateGuard
guard = _MintRateGuard("http://mint:3338", 2)
guard._cooldown_until = 15.0
operation = AsyncMock(return_value="ok")
with patch("routstr.wallet.time.monotonic", return_value=10.0):
with patch("routstr.wallet.asyncio.sleep", AsyncMock()) as sleep:
assert await guard.run(operation) == "ok"
sleep.assert_awaited_once_with(5.0)
operation.assert_awaited_once()
def test_mint_rate_guard_rebuilds_when_setting_changes() -> None:
from routstr.core.settings import settings
from routstr.wallet import _MintRateGuard
with patch.object(settings, "mint_max_concurrency", 4):
first = _MintRateGuard.get("http://mint:3338")
with patch.object(settings, "mint_max_concurrency", 2):
second = _MintRateGuard.get("http://mint:3338")
assert first is not None
assert second is not None
assert first is not second
assert second._max_concurrency == 2
@pytest.mark.asyncio
async def test_mint_operation_honors_retry_after_as_minimum() -> None:
from routstr.core.settings import settings
from routstr.wallet import _mint_operation
request = httpx.Request("POST", "http://mint:3338/v1/mint/quote/bolt11")
response = httpx.Response(429, request=request, headers={"Retry-After": "60"})
calls = 0
async def factory() -> str:
nonlocal calls
calls += 1
if calls == 1:
raise httpx.HTTPStatusError(
"rate limited", request=request, response=response
)
return "ok"
sleep = AsyncMock()
with patch.object(settings, "mint_retry_max_attempts", 1):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch.object(settings, "mint_max_concurrency", 1):
with patch("routstr.wallet.time.monotonic", return_value=0.1):
with patch("routstr.wallet.asyncio.sleep", sleep):
result = await _mint_operation(
factory, mint_url="http://mint:3338"
)
assert result == "ok"
sleep.assert_awaited_once_with(60.0)
@pytest.mark.asyncio
async def test_mint_operation_timeout_includes_adaptive_cooldown() -> None:
from routstr.core.settings import settings
from routstr.wallet import _mint_operation, _MintRateGuard
operation = AsyncMock(return_value="unexpected")
with patch.object(settings, "mint_max_concurrency", 1):
guard = _MintRateGuard.get("http://mint:3338")
assert guard is not None
guard.apply_cooldown(60)
with patch.object(settings, "mint_operation_timeout_seconds", 0.01):
with pytest.raises(httpx.TimeoutException, match="total timeout"):
await _mint_operation(operation, mint_url="http://mint:3338")
operation.assert_not_awaited()
@pytest.mark.asyncio
async def test_mint_operation_retries_httpx_timeout_only_when_safe() -> None:
from routstr.core.settings import settings
from routstr.wallet import _mint_operation
retrying = AsyncMock(side_effect=[httpx.ReadTimeout("slow"), "ok"])
non_retrying = AsyncMock(side_effect=httpx.ReadTimeout("ambiguous"))
with patch.object(settings, "mint_retry_max_attempts", 2):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch("routstr.wallet.asyncio.sleep", AsyncMock()):
assert await _mint_operation(retrying) == "ok"
with pytest.raises(httpx.TimeoutException):
await _mint_operation(non_retrying, retry_timeouts=False)
assert retrying.await_count == 2
assert non_retrying.await_count == 1
@pytest.mark.asyncio
async def test_get_wallet_initializes_and_loads_once_concurrently() -> None:
from routstr.wallet import get_wallet
mock_wallet = Mock()
mock_wallet.load_mint = AsyncMock()
mock_wallet.load_proofs = AsyncMock()
with patch(
"routstr.wallet.Wallet.with_db", AsyncMock(return_value=mock_wallet)
) as create:
with patch("routstr.wallet.time.monotonic", return_value=100.0):
first, second = await asyncio.gather(
get_wallet("http://mint:3338"), get_wallet("http://mint:3338")
)
assert first is second is mock_wallet
create.assert_awaited_once()
mock_wallet.load_mint.assert_awaited_once()
mock_wallet.load_proofs.assert_awaited_once_with(reload=True)
@pytest.mark.asyncio
async def test_mint_operation_factory_retry_succeeds() -> None:
"""_mint_operation accepts a zero-arg factory, not a dead coroutine.
A factory that raises twice then succeeds must be retried and return."""
from routstr.core.settings import settings
from routstr.wallet import _mint_operation
calls = 0
async def factory() -> str:
nonlocal calls
calls += 1
if calls < 3:
raise TimeoutError("timeout")
return "ok"
with patch.object(settings, "mint_retry_max_attempts", 3):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch.object(settings, "mint_max_concurrency", 0):
with patch("asyncio.sleep", AsyncMock()):
result = await _mint_operation(
factory, op_name="test_retry", mint_url="http://mint:3338"
)
assert calls == 3
assert result == "ok"
@pytest.mark.asyncio
async def test_mint_operation_factory_retry_exhausted() -> None:
"""When the factory always times out, _mint_operation raises
httpx.TimeoutException after mint_retry_max_attempts + 1 attempts."""
from routstr.core.settings import settings
from routstr.wallet import _mint_operation
calls = 0
async def factory() -> None:
nonlocal calls
calls += 1
raise TimeoutError("always timeout")
with patch.object(settings, "mint_retry_max_attempts", 2):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch.object(settings, "mint_max_concurrency", 0):
with patch("asyncio.sleep", AsyncMock()):
with pytest.raises(httpx.TimeoutException):
await _mint_operation(
factory, op_name="test_exhaust", mint_url="http://mint:3338"
)
assert calls == 3 # max_attempts(2) + 1
# ---------------------------------------------------------------------------
# Trusted-mint fallback
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_lightning_mint_fallback_for_topups() -> None:
"""When the primary mint is unreachable, _request_mint_with_fallback
falls back to a secondary trusted mint."""
from routstr.core.settings import settings
from routstr.lightning import _request_mint_with_fallback
primary = "http://primary:3338"
secondary = "http://secondary:3338"
mock_primary_wallet = Mock()
mock_primary_wallet.request_mint = AsyncMock(
side_effect=httpx.ConnectError("primary down")
)
mock_quote = Mock()
mock_quote.request = "lnbc1secondary"
mock_quote.quote = "quote_secondary"
mock_secondary_wallet = Mock()
mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote)
wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet}
mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m])
with patch.object(settings, "primary_mint", primary):
with patch.object(settings, "cashu_mints", [primary, secondary]):
with patch.object(settings, "mint_max_concurrency", 0):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch("routstr.lightning.get_wallet", side_effect=mock_get):
bolt11, quote_id, mint_url = await _request_mint_with_fallback(
1000
)
assert mint_url == secondary
assert bolt11 == "lnbc1secondary"
assert quote_id == "quote_secondary"
mock_primary_wallet.request_mint.assert_called_once()
mock_secondary_wallet.request_mint.assert_called_once()
@pytest.mark.asyncio
async def test_swap_falls_back_to_secondary_mint() -> None:
"""When the primary mint is unreachable, swap_to_primary_mint falls back
to a secondary trusted mint as the swap destination."""
from routstr.core.settings import settings
from routstr.wallet import _wallet_last_load, _wallets, swap_to_primary_mint
_wallets.clear()
_wallet_last_load.clear()
primary = "http://primary:3338"
secondary = "http://secondary:3338"
foreign = "http://foreign:3338"
mock_token = Mock()
mock_token.mint = foreign
mock_token.unit = "sat"
mock_token.amount = 1000
mock_token.keysets = ["keyset1"]
mock_token.proofs = [Mock(amount=1000)]
mock_token_wallet = Mock()
mock_token_wallet.load_mint = AsyncMock()
mock_token_wallet.load_proofs = AsyncMock()
mock_token_wallet.get_fees_for_proofs = Mock(return_value=0)
mock_token_wallet.melt_quote = AsyncMock(
return_value=Mock(quote="melt_q", amount=990, fee_reserve=10)
)
mock_token_wallet.melt = AsyncMock(return_value=Mock())
mock_primary_wallet = Mock()
mock_primary_wallet.request_mint = AsyncMock(
side_effect=httpx.ConnectError("primary down")
)
mint_quote = Mock(quote="mint_q_secondary", request="lnbc1secondary")
mock_secondary_wallet = Mock()
mock_secondary_wallet.load_mint = AsyncMock()
mock_secondary_wallet.load_proofs = AsyncMock()
mock_secondary_wallet.available_balance = Mock(amount=0)
mock_secondary_wallet.keysets = ["ks_secondary"]
mock_secondary_wallet.restore_tokens_for_keyset = AsyncMock()
mock_secondary_wallet.request_mint = AsyncMock(return_value=mint_quote)
mock_secondary_wallet.mint = AsyncMock(return_value=Mock())
wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet}
mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m])
with patch.object(settings, "primary_mint", primary):
with patch.object(settings, "primary_mint_unit", "sat"):
with patch.object(settings, "cashu_mints", [primary, secondary]):
with patch.object(settings, "mint_max_concurrency", 0):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch("asyncio.sleep", AsyncMock()):
with patch(
"routstr.wallet.get_wallet", side_effect=mock_get
):
amount, unit, mint_url = await swap_to_primary_mint(
mock_token, mock_token_wallet
)
assert mint_url == secondary
assert amount == 990 # 1000 - 10 fee_reserve
assert unit == "sat"
mock_secondary_wallet.mint.assert_called_once()
mock_primary_wallet.mint.assert_not_called()
@pytest.mark.asyncio
async def test_lightning_mint_fallback_on_429() -> None:
"""A 429 from the primary mint should trigger fallback to a secondary,
not just transport errors."""
from routstr.core.settings import settings
from routstr.lightning import _request_mint_with_fallback
primary = "http://primary:3338"
secondary = "http://secondary:3338"
mock_resp = Mock(status_code=429, headers={})
mock_resp.raise_for_status = Mock(
side_effect=httpx.HTTPStatusError(
"rate limited", request=Mock(), response=mock_resp
)
)
mock_primary_wallet = Mock()
mock_primary_wallet.request_mint = AsyncMock(
side_effect=httpx.HTTPStatusError(
"rate limited", request=Mock(), response=mock_resp
)
)
mock_quote = Mock(request="lnbc1secondary", quote="quote_secondary")
mock_secondary_wallet = Mock()
mock_secondary_wallet.request_mint = AsyncMock(return_value=mock_quote)
wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet}
mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m])
with patch.object(settings, "primary_mint", primary):
with patch.object(settings, "cashu_mints", [primary, secondary]):
with patch.object(settings, "mint_retry_max_attempts", 0):
with patch.object(settings, "mint_max_concurrency", 0):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch(
"routstr.lightning.get_wallet", side_effect=mock_get
):
(
bolt11,
quote_id,
mint_url,
) = await _request_mint_with_fallback(1000)
assert mint_url == secondary
mock_secondary_wallet.request_mint.assert_called_once()
@pytest.mark.asyncio
async def test_lightning_mint_fallback_all_fail() -> None:
"""When every trusted mint fails, _request_mint_with_fallback raises
MintConnectionError instead of trying indefinitely."""
from routstr.core.settings import settings
from routstr.lightning import _request_mint_with_fallback
from routstr.wallet import MintConnectionError
primary = "http://primary:3338"
secondary = "http://secondary:3338"
mock_primary_wallet = Mock()
mock_primary_wallet.request_mint = AsyncMock(side_effect=httpx.ConnectError("down"))
mock_secondary_wallet = Mock()
mock_secondary_wallet.request_mint = AsyncMock(
side_effect=httpx.ConnectError("down")
)
wallets_map = {primary: mock_primary_wallet, secondary: mock_secondary_wallet}
mock_get = AsyncMock(side_effect=lambda m, *a, **kw: wallets_map[m])
with patch.object(settings, "primary_mint", primary):
with patch.object(settings, "cashu_mints", [primary, secondary]):
with patch.object(settings, "mint_retry_max_attempts", 0):
with patch.object(settings, "mint_max_concurrency", 0):
with patch.object(settings, "mint_operation_timeout_seconds", 0):
with patch(
"routstr.lightning.get_wallet", side_effect=mock_get
):
with pytest.raises(MintConnectionError):
await _request_mint_with_fallback(1000)