mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-30 15:26:14 +00:00
Compare commits
50
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
443c910b9e | ||
|
|
88d301398b | ||
|
|
4cc9aef61f | ||
|
|
3befe063f4 | ||
|
|
f8adaee362 | ||
|
|
895ea90bfa | ||
|
|
c4d27ba02a | ||
|
|
5ea5024608 | ||
|
|
c75dee147a | ||
|
|
48c11eb7bc | ||
|
|
c829685f80 | ||
|
|
39f801561b | ||
|
|
ff55788e2d | ||
|
|
1138cdd4ef | ||
|
|
f15eab9f10 | ||
|
|
1131c2d583 | ||
|
|
7108d554c8 | ||
|
|
1eddf89d52 | ||
|
|
a2cedd6769 | ||
|
|
7b0ade3987 | ||
|
|
2410a4a6ce | ||
|
|
1b09639265 | ||
|
|
7394b10e75 | ||
|
|
92246b78d0 | ||
|
|
6023c03959 | ||
|
|
a9a6381614 | ||
|
|
040799a4d7 | ||
|
|
586af15a1b | ||
|
|
3e906605a0 | ||
|
|
1957e716a3 | ||
|
|
69f19ff991 | ||
|
|
8b942f3c14 | ||
|
|
09e1c7bf2d | ||
|
|
cc2a96e2ef | ||
|
|
93ab1d927b | ||
|
|
39970d8bee | ||
|
|
d7c401d204 | ||
|
|
d44b98fd0d | ||
|
|
6fa3610423 | ||
|
|
65702171e4 | ||
|
|
65abcbce92 | ||
|
|
40153d4c36 | ||
|
|
40bf976fbc | ||
|
|
eae20f04a7 | ||
|
|
acb630f6cf | ||
|
|
1230d528de | ||
|
|
d23c90b939 | ||
|
|
d8db2a3051 | ||
|
|
0bbbf902cd | ||
|
|
7ed18a9d02 |
@@ -22,6 +22,19 @@ ROUTSTR_SECRET_KEY=
|
||||
|
||||
# Database
|
||||
# DATABASE_URL=sqlite+aiosqlite:///keys.db
|
||||
# Pool controls are validated at boot, sourced only from the environment, and
|
||||
# logged at startup. Keep total capacity across all workers below the database
|
||||
# connection limit. Pre-ping is automatic for networked backends; SQLite may
|
||||
# explicitly opt in if desired.
|
||||
# DATABASE_POOL_SIZE=5
|
||||
# DATABASE_MAX_OVERFLOW=10
|
||||
# DATABASE_POOL_TIMEOUT=30
|
||||
# DATABASE_POOL_RECYCLE=1800
|
||||
# DATABASE_POOL_PRE_PING=false
|
||||
# Warn when a checkout is held this many seconds.
|
||||
# DATABASE_POOL_HOLD_WARN_SECONDS=10
|
||||
# SQLite serialises writes; increasing its pool can trade pool timeouts for
|
||||
# "database is locked" errors rather than increasing write throughput.
|
||||
|
||||
# Node Information
|
||||
# NAME=My Routstr Node
|
||||
@@ -31,7 +44,9 @@ ROUTSTR_SECRET_KEY=
|
||||
# RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com"
|
||||
# ENABLE_ANALYTICS_SHARING=true
|
||||
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me"
|
||||
# MINT_OPERATION_CONCURRENCY=4
|
||||
# RECEIVE_LN_ADDRESS=
|
||||
# REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS=900
|
||||
|
||||
# Custom Pricing Configuration
|
||||
# MODEL_BASED_PRICING=true
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
"""add refund sweep claim lease
|
||||
|
||||
Revision ID: aa50fde387a2
|
||||
Revises: 9c4d8e2f1a6b
|
||||
Create Date: 2026-07-26 12:50:10.509217
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "aa50fde387a2"
|
||||
down_revision = "9c4d8e2f1a6b"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
columns = {
|
||||
column["name"]
|
||||
for column in sa.inspect(conn).get_columns("cashu_transactions")
|
||||
}
|
||||
if "sweep_started_at" not in columns:
|
||||
op.add_column(
|
||||
"cashu_transactions",
|
||||
sa.Column("sweep_started_at", sa.Integer(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
columns = {
|
||||
column["name"]
|
||||
for column in sa.inspect(conn).get_columns("cashu_transactions")
|
||||
}
|
||||
if "sweep_started_at" in columns:
|
||||
op.drop_column("cashu_transactions", "sweep_started_at")
|
||||
@@ -0,0 +1,25 @@
|
||||
"""add mint url to lightning invoices
|
||||
|
||||
Revision ID: bf76270b66c4
|
||||
Revises: aa50fde387a2
|
||||
Create Date: 2026-07-30 00:54:30.306876
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "bf76270b66c4"
|
||||
down_revision = "aa50fde387a2"
|
||||
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")
|
||||
+46
-14
@@ -33,6 +33,7 @@ from .wallet import (
|
||||
classify_redemption_error,
|
||||
credit_balance,
|
||||
deserialize_token_from_string,
|
||||
wallet_operation_guard,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -50,6 +51,24 @@ ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
|
||||
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
|
||||
|
||||
|
||||
def _format_msat_amount(amount: int) -> str:
|
||||
sats = f"{amount / 1000:.3f}".rstrip("0").rstrip(".")
|
||||
return f"{sats} sats ({amount} msats)"
|
||||
|
||||
|
||||
def _model_balance_error(required: int, available: int) -> dict[str, dict[str, str]]:
|
||||
return {
|
||||
"error": {
|
||||
"message": (
|
||||
f"Insufficient balance: {_format_msat_amount(required)} required "
|
||||
f"for this model; {_format_msat_amount(available)} available."
|
||||
),
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance",
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ReservationSnapshot:
|
||||
release_id: str
|
||||
@@ -153,6 +172,29 @@ async def validate_bearer_key(
|
||||
refund_address: Optional[str] = None,
|
||||
key_expiry_time: Optional[int] = None,
|
||||
min_cost: int = 0,
|
||||
) -> ApiKey:
|
||||
if bearer_key.startswith("cashu"):
|
||||
# Acquire before the first lookup/flush so concurrent token creation
|
||||
# cannot hold SQLite write transactions while waiting to mutate proofs.
|
||||
async with wallet_operation_guard():
|
||||
return await _validate_bearer_key_locked(
|
||||
bearer_key,
|
||||
session,
|
||||
refund_address,
|
||||
key_expiry_time,
|
||||
min_cost,
|
||||
)
|
||||
return await _validate_bearer_key_locked(
|
||||
bearer_key, session, refund_address, key_expiry_time, min_cost
|
||||
)
|
||||
|
||||
|
||||
async def _validate_bearer_key_locked(
|
||||
bearer_key: str,
|
||||
session: AsyncSession,
|
||||
refund_address: Optional[str] = None,
|
||||
key_expiry_time: Optional[int] = None,
|
||||
min_cost: int = 0,
|
||||
) -> ApiKey:
|
||||
"""
|
||||
Validates the provided API key using SQLModel.
|
||||
@@ -240,13 +282,7 @@ async def validate_bearer_key(
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Insufficient balance: {min_cost} mSats required for this model. {billing_key.total_balance} available.",
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance",
|
||||
}
|
||||
},
|
||||
detail=_model_balance_error(min_cost, billing_key.total_balance),
|
||||
)
|
||||
|
||||
# Early check: Spending limit check (Child key limit)
|
||||
@@ -336,13 +372,9 @@ async def validate_bearer_key(
|
||||
if min_cost > 0 and existing_key.total_balance < min_cost:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"error": {
|
||||
"message": f"Insufficient balance: {min_cost} mSats required for this model. {existing_key.total_balance} available.",
|
||||
"type": "insufficient_quota",
|
||||
"code": "insufficient_balance",
|
||||
}
|
||||
},
|
||||
detail=_model_balance_error(
|
||||
min_cost, existing_key.total_balance
|
||||
),
|
||||
)
|
||||
|
||||
return existing_key
|
||||
|
||||
+84
-13
@@ -30,6 +30,7 @@ from .wallet import (
|
||||
recieve_token,
|
||||
send_to_lnurl,
|
||||
send_token,
|
||||
token_mint_url,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -184,6 +185,17 @@ class TopupRequest(BaseModel):
|
||||
cashu_token: str
|
||||
|
||||
|
||||
def _error_chain(error: BaseException) -> list[dict[str, str]]:
|
||||
chain: list[dict[str, str]] = []
|
||||
current: BaseException | None = error
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
chain.append({"type": type(current).__name__, "message": str(current)})
|
||||
current = current.__cause__ or current.__context__
|
||||
return chain
|
||||
|
||||
|
||||
@router.post("/topup")
|
||||
async def topup_wallet_endpoint(
|
||||
cashu_token: str | None = None,
|
||||
@@ -201,6 +213,18 @@ async def topup_wallet_endpoint(
|
||||
cashu_token = cashu_token.replace("\n", "").replace("\r", "").replace("\t", "")
|
||||
if len(cashu_token) < 10 or "cashu" not in cashu_token:
|
||||
raise HTTPException(status_code=400, detail="Invalid token format")
|
||||
|
||||
source_mint = token_mint_url(cashu_token, "unknown")
|
||||
logger.warning(
|
||||
"Cashu wallet top-up started",
|
||||
extra={
|
||||
"event": "cashu_topup_started",
|
||||
"source_mint": source_mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"trusted_mints": settings.cashu_mints,
|
||||
"key_hash": billing_key.hashed_key[:8],
|
||||
},
|
||||
)
|
||||
try:
|
||||
amount_msats = await credit_balance(cashu_token, billing_key, session)
|
||||
except Exception as e:
|
||||
@@ -209,12 +233,41 @@ async def topup_wallet_endpoint(
|
||||
classified = classify_redemption_error(e)
|
||||
if classified is None:
|
||||
logger.error(
|
||||
"topup_wallet_endpoint: unhandled error",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
"Cashu wallet top-up failed with an unhandled error",
|
||||
extra={
|
||||
"event": "cashu_topup_failed",
|
||||
"source_mint": source_mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"trusted_mints": settings.cashu_mints,
|
||||
"error_chain": _error_chain(e),
|
||||
},
|
||||
)
|
||||
raise HTTPException(status_code=500, detail="Internal server error")
|
||||
_type, status_code, message, _code = classified
|
||||
error_type, status_code, message, error_code = classified
|
||||
logger.warning(
|
||||
"Cashu wallet top-up failed",
|
||||
extra={
|
||||
"event": "cashu_topup_failed",
|
||||
"source_mint": source_mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"trusted_mints": settings.cashu_mints,
|
||||
"status_code": status_code,
|
||||
"error_type": error_type,
|
||||
"error_code": error_code,
|
||||
"error_chain": _error_chain(e),
|
||||
},
|
||||
)
|
||||
raise HTTPException(status_code=status_code, detail=message)
|
||||
|
||||
logger.warning(
|
||||
"Cashu wallet top-up completed",
|
||||
extra={
|
||||
"event": "cashu_topup_completed",
|
||||
"source_mint": source_mint,
|
||||
"credited_msats": amount_msats,
|
||||
"key_hash": billing_key.hashed_key[:8],
|
||||
},
|
||||
)
|
||||
return {"msats": amount_msats}
|
||||
|
||||
|
||||
@@ -260,7 +313,11 @@ async def _lookup_key_no_create(
|
||||
|
||||
|
||||
async def _restore_balance(
|
||||
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
|
||||
session: AsyncSession,
|
||||
hashed_key: str,
|
||||
balance: int,
|
||||
reserved_balance: int,
|
||||
mint_url: str,
|
||||
) -> None:
|
||||
"""Restore balance after a failed refund mint attempt."""
|
||||
restore_stmt = (
|
||||
@@ -275,7 +332,11 @@ async def _restore_balance(
|
||||
await session.commit()
|
||||
logger.info(
|
||||
"refund_wallet_endpoint: balance restored after mint failure",
|
||||
extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url},
|
||||
extra={
|
||||
"hashed_key": hashed_key,
|
||||
"restored_balance": balance,
|
||||
"mint_url": mint_url,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -418,15 +479,14 @@ async def refund_wallet_endpoint(
|
||||
detail="Balance changed concurrently. Please retry the refund.",
|
||||
)
|
||||
|
||||
# --- MINT: balance is locked at zero, safe to create the refund token ---
|
||||
# Proofs from untrusted mints are swapped to primary_mint on receive.
|
||||
# Use primary_mint unless key.refund_mint_url is an explicitly trusted mint.
|
||||
# The balance is locked at zero, so it is safe to create the refund token.
|
||||
effective_refund_mint = (
|
||||
key.refund_mint_url
|
||||
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints
|
||||
else settings.primary_mint
|
||||
)
|
||||
try:
|
||||
refund_currency = key.refund_currency or "sat"
|
||||
if key.refund_address:
|
||||
await send_to_lnurl(
|
||||
remaining_balance,
|
||||
@@ -436,10 +496,10 @@ async def refund_wallet_endpoint(
|
||||
)
|
||||
result = {"recipient": key.refund_address}
|
||||
else:
|
||||
refund_currency = key.refund_currency or "sat"
|
||||
token = await send_token(
|
||||
remaining_balance, refund_currency, effective_refund_mint
|
||||
)
|
||||
effective_refund_mint = token_mint_url(token, effective_refund_mint)
|
||||
result = {"token": token}
|
||||
|
||||
if key.refund_currency == "sat":
|
||||
@@ -460,11 +520,23 @@ async def refund_wallet_endpoint(
|
||||
|
||||
except HTTPException:
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
await _restore_balance(
|
||||
session,
|
||||
key.hashed_key,
|
||||
pre_debit_balance,
|
||||
pre_debit_reserved,
|
||||
key.refund_mint_url or "",
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
await _restore_balance(
|
||||
session,
|
||||
key.hashed_key,
|
||||
pre_debit_balance,
|
||||
pre_debit_reserved,
|
||||
key.refund_mint_url or "",
|
||||
)
|
||||
error_msg = str(e)
|
||||
logger.error(
|
||||
"refund_wallet_endpoint: mint/send failed",
|
||||
@@ -491,7 +563,7 @@ async def refund_wallet_endpoint(
|
||||
token=result["token"],
|
||||
amount=remaining_balance,
|
||||
unit=key.refund_currency or "sat",
|
||||
mint_url=key.refund_mint_url,
|
||||
mint_url=effective_refund_mint,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="apikey",
|
||||
@@ -685,7 +757,6 @@ async def reset_child_key_spent(
|
||||
return {"success": True, "message": "Child key balance reset successfully."}
|
||||
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/{path:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE"],
|
||||
|
||||
@@ -1669,12 +1669,18 @@ async def get_transactions_api(
|
||||
)
|
||||
total = count_result.one()
|
||||
|
||||
stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit)
|
||||
stmt = (
|
||||
base.order_by(col(CashuTransaction.created_at).desc())
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
results = await session.exec(stmt)
|
||||
transactions = results.all()
|
||||
|
||||
return {
|
||||
"transactions": [tx.dict() for tx in transactions],
|
||||
"transactions": [
|
||||
tx.dict(exclude={"sweep_started_at"}) for tx in transactions
|
||||
],
|
||||
"total": total,
|
||||
}
|
||||
|
||||
|
||||
+113
-8
@@ -12,21 +12,80 @@ from typing import AsyncGenerator
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from alembic.util.exc import CommandError
|
||||
from sqlalchemy import Index, UniqueConstraint, case, delete, or_
|
||||
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_
|
||||
from sqlalchemy.engine import make_url
|
||||
from sqlalchemy.exc import IntegrityError, OperationalError
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
||||
from sqlalchemy.orm import aliased
|
||||
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .logging import get_logger
|
||||
from .settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
|
||||
|
||||
|
||||
engine = create_async_engine(DATABASE_URL, echo=False) # echo=True for debugging SQL
|
||||
def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
|
||||
"""Build and instrument an async engine from environment-only settings."""
|
||||
url = make_url(database_url)
|
||||
backend = url.get_backend_name()
|
||||
is_sqlite = backend == "sqlite"
|
||||
is_memory_sqlite = is_sqlite and url.database in {None, "", ":memory:"}
|
||||
pool_pre_ping = settings.database_pool_pre_ping or not is_sqlite
|
||||
options: dict[str, int | float | bool] = {"pool_pre_ping": pool_pre_ping}
|
||||
if not is_memory_sqlite:
|
||||
options.update(
|
||||
pool_size=settings.database_pool_size,
|
||||
max_overflow=settings.database_max_overflow,
|
||||
pool_timeout=settings.database_pool_timeout,
|
||||
pool_recycle=settings.database_pool_recycle,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Database pool configured",
|
||||
extra={
|
||||
"database_url_backend": backend,
|
||||
"in_memory_sqlite": is_memory_sqlite,
|
||||
**options,
|
||||
},
|
||||
)
|
||||
created_engine = create_async_engine(database_url, echo=False, **options)
|
||||
hold_warn_seconds = settings.database_pool_hold_warn_seconds
|
||||
|
||||
def record_pool_checkout(
|
||||
dbapi_connection: object, connection_record: object, proxy: object
|
||||
) -> None:
|
||||
connection_record.info["routstr_checked_out_at"] = time.monotonic() # type: ignore[attr-defined]
|
||||
|
||||
def record_pool_checkin(
|
||||
dbapi_connection: object, connection_record: object
|
||||
) -> None:
|
||||
checked_out_at = connection_record.info.pop( # type: ignore[attr-defined]
|
||||
"routstr_checked_out_at", None
|
||||
)
|
||||
if checked_out_at is None:
|
||||
return
|
||||
held_seconds = time.monotonic() - checked_out_at
|
||||
if held_seconds >= hold_warn_seconds:
|
||||
logger.warning(
|
||||
"Database connection held longer than threshold",
|
||||
extra={
|
||||
"held_seconds": round(held_seconds, 3),
|
||||
"threshold_seconds": hold_warn_seconds,
|
||||
"pool_status": created_engine.pool.status(),
|
||||
},
|
||||
)
|
||||
|
||||
event.listen(created_engine.sync_engine, "checkout", record_pool_checkout)
|
||||
event.listen(created_engine.sync_engine, "checkin", record_pool_checkin)
|
||||
return created_engine
|
||||
|
||||
|
||||
engine = create_db_engine()
|
||||
|
||||
|
||||
class ApiKey(SQLModel, table=True): # type: ignore
|
||||
@@ -222,7 +281,10 @@ async def release_stale_reservations(
|
||||
if released:
|
||||
logger.warning(
|
||||
"Released stale reservations",
|
||||
extra={"released_reservations": released, "max_age_seconds": max_age_seconds},
|
||||
extra={
|
||||
"released_reservations": released,
|
||||
"max_age_seconds": max_age_seconds,
|
||||
},
|
||||
)
|
||||
return released
|
||||
|
||||
@@ -320,12 +382,16 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
|
||||
description: str = Field(description="Invoice description")
|
||||
payment_hash: str = Field(description="Payment hash for tracking", unique=True)
|
||||
status: str = Field(
|
||||
default="pending", description="pending, paid, expired, cancelled"
|
||||
default="pending",
|
||||
description="pending, paid, expired, cancelled, reconciliation_required",
|
||||
)
|
||||
api_key_hash: str | None = Field(
|
||||
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"
|
||||
)
|
||||
@@ -365,6 +431,10 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
collected: bool = Field(default=False)
|
||||
swept: bool = Field(default=False)
|
||||
sweep_started_at: int | None = Field(
|
||||
default=None,
|
||||
description="Unix timestamp for a recoverable refund-sweep claim",
|
||||
)
|
||||
source: str = Field(
|
||||
default="x-cashu",
|
||||
description="Payment source: x-cashu or apikey",
|
||||
@@ -717,14 +787,49 @@ async def complete_routstr_fee_payout(
|
||||
return result.rowcount == 1
|
||||
|
||||
|
||||
async def balances_for_mint_and_unit(
|
||||
async def total_user_liability(db_session: AsyncSession) -> int:
|
||||
"""Return all outstanding API-key balances in millisatoshis."""
|
||||
result = await db_session.exec(select(func.sum(ApiKey.balance)))
|
||||
return int(result.one() or 0)
|
||||
|
||||
|
||||
async def balance_for_mint_and_unit(
|
||||
db_session: AsyncSession, mint_url: str, unit: str
|
||||
) -> int:
|
||||
query = select(func.sum(ApiKey.balance)).where(
|
||||
ApiKey.refund_mint_url == mint_url, ApiKey.refund_currency == unit
|
||||
"""Return the user liability for one mint and unit in millisatoshis."""
|
||||
result = await db_session.exec(
|
||||
select(func.sum(ApiKey.balance)).where(
|
||||
col(ApiKey.refund_mint_url) == mint_url,
|
||||
col(ApiKey.refund_currency) == unit,
|
||||
)
|
||||
)
|
||||
return int(result.one() or 0)
|
||||
|
||||
|
||||
async def balances_by_mint_and_unit(
|
||||
db_session: AsyncSession, mint_urls: list[str], units: list[str]
|
||||
) -> dict[tuple[str, str], int]:
|
||||
"""Return requested user liabilities grouped by mint and unit."""
|
||||
if not mint_urls or not units:
|
||||
return {}
|
||||
query = (
|
||||
select(
|
||||
col(ApiKey.refund_mint_url),
|
||||
col(ApiKey.refund_currency),
|
||||
func.sum(ApiKey.balance),
|
||||
)
|
||||
.where(
|
||||
col(ApiKey.refund_mint_url).in_(mint_urls),
|
||||
col(ApiKey.refund_currency).in_(units),
|
||||
)
|
||||
.group_by(col(ApiKey.refund_mint_url), col(ApiKey.refund_currency))
|
||||
)
|
||||
result = await db_session.exec(query)
|
||||
return result.one() or 0
|
||||
return {
|
||||
(mint_url, unit): int(balance or 0)
|
||||
for mint_url, unit, balance in result.all()
|
||||
if mint_url is not None and unit is not None
|
||||
}
|
||||
|
||||
|
||||
async def init_db() -> None:
|
||||
|
||||
@@ -41,6 +41,9 @@ class Settings(BaseSettings):
|
||||
receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS")
|
||||
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
|
||||
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
|
||||
mint_operation_concurrency: int = Field(
|
||||
default=4, ge=1, env="MINT_OPERATION_CONCURRENCY"
|
||||
)
|
||||
|
||||
# Lightning payout configuration
|
||||
# Minimum available balance (in satoshis) before profit is paid out over
|
||||
@@ -50,6 +53,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
|
||||
@@ -98,7 +113,27 @@ 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"
|
||||
)
|
||||
refund_sweep_claim_timeout_seconds: int = Field(
|
||||
default=900, gt=0, env="REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS"
|
||||
)
|
||||
|
||||
# Database connection-pool controls (advanced). Capacity defaults match
|
||||
# SQLAlchemy's established queue-pool behavior. Pre-ping is enabled by the
|
||||
# engine factory for networked backends; SQLite can explicitly opt in.
|
||||
# These fields are env-only below.
|
||||
database_pool_size: int = Field(default=5, ge=1, env="DATABASE_POOL_SIZE")
|
||||
database_max_overflow: int = Field(default=10, ge=0, env="DATABASE_MAX_OVERFLOW")
|
||||
database_pool_timeout: float = Field(
|
||||
default=30.0, gt=0, env="DATABASE_POOL_TIMEOUT"
|
||||
)
|
||||
database_pool_recycle: int = Field(default=1800, ge=0, env="DATABASE_POOL_RECYCLE")
|
||||
database_pool_pre_ping: bool = Field(default=False, env="DATABASE_POOL_PRE_PING")
|
||||
database_pool_hold_warn_seconds: float = Field(
|
||||
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
|
||||
)
|
||||
|
||||
# Logging
|
||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||
@@ -117,9 +152,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."""
|
||||
@@ -144,10 +178,32 @@ def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
|
||||
# ``routstr.core.vault``.
|
||||
SECRET_FIELDS = frozenset({"admin_password", "nsec"})
|
||||
|
||||
# Infrastructure the node needs *before* it can open a DB session — so it can
|
||||
# never be configured from the DB (chicken-and-egg) and stays env-only. Unlike
|
||||
# secrets (owned by bootstrap), these are excluded so the DB settings blob can
|
||||
# neither store nor shadow them; env is always authoritative.
|
||||
ENV_ONLY_FIELDS = frozenset(
|
||||
{
|
||||
"database_pool_size",
|
||||
"database_max_overflow",
|
||||
"database_pool_timeout",
|
||||
"database_pool_recycle",
|
||||
"database_pool_pre_ping",
|
||||
"database_pool_hold_warn_seconds",
|
||||
}
|
||||
)
|
||||
|
||||
_NON_PERSISTED_FIELDS = SECRET_FIELDS | ENV_ONLY_FIELDS
|
||||
|
||||
|
||||
def _strip_secret_fields(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Return a copy of ``data`` without any secret fields (for persistence)."""
|
||||
return {k: v for k, v in data.items() if k not in SECRET_FIELDS}
|
||||
"""Return a copy of ``data`` without secret or env-only fields.
|
||||
|
||||
Both are kept out of the persisted settings blob: secrets for confidentiality,
|
||||
env-only fields (e.g. DB pool sizing) because they must never be sourced from
|
||||
the database.
|
||||
"""
|
||||
return {k: v for k, v in data.items() if k not in _NON_PERSISTED_FIELDS}
|
||||
|
||||
|
||||
def _apply_to_live_settings(data: dict[str, Any]) -> None:
|
||||
@@ -327,7 +383,13 @@ 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
|
||||
and k not in ENV_ONLY_FIELDS
|
||||
}
|
||||
)
|
||||
merged_dict = Settings(**merged_dict).dict()
|
||||
|
||||
@@ -394,8 +456,13 @@ class SettingsService:
|
||||
)
|
||||
)
|
||||
await db_session.commit()
|
||||
# Update in-place
|
||||
# Update in-place. Env-only fields (e.g. DB pool sizing) are never
|
||||
# applied here: the engine pool is already built at boot from env,
|
||||
# so letting an update mutate the live value would only make it
|
||||
# diverge from the running pool.
|
||||
for k, v in candidate.dict().items():
|
||||
if k in ENV_ONLY_FIELDS:
|
||||
continue
|
||||
setattr(settings, k, v)
|
||||
cls._current = settings
|
||||
return settings
|
||||
|
||||
+364
-56
@@ -1,22 +1,73 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import re
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlmodel import col, select
|
||||
from sqlalchemy.orm.attributes import set_committed_value
|
||||
from sqlmodel import col, select, update
|
||||
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_cooldown_remaining,
|
||||
_mint_operation,
|
||||
get_wallet,
|
||||
is_mint_connection_error,
|
||||
wallet_operation_guard,
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
lightning_router = APIRouter(prefix="/lightning")
|
||||
|
||||
# Avoid duplicate work within one process. Cross-process credit fencing is done
|
||||
# by the conditional pending -> paid update in _finalize_invoice_settlement().
|
||||
_invoice_settlement_locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _InvoiceSettlement:
|
||||
id: str
|
||||
payment_hash: str
|
||||
amount_sats: int
|
||||
purpose: str
|
||||
api_key_hash: str | None
|
||||
mint_url: str | None
|
||||
balance_limit: int | None
|
||||
balance_limit_reset: str | None
|
||||
validity_date: int | None
|
||||
|
||||
@classmethod
|
||||
def from_invoice(cls, invoice: LightningInvoice) -> "_InvoiceSettlement":
|
||||
return cls(
|
||||
id=invoice.id,
|
||||
payment_hash=invoice.payment_hash,
|
||||
amount_sats=invoice.amount_sats,
|
||||
purpose=invoice.purpose,
|
||||
api_key_hash=invoice.api_key_hash,
|
||||
mint_url=invoice.mint_url,
|
||||
balance_limit=invoice.balance_limit,
|
||||
balance_limit_reset=invoice.balance_limit_reset,
|
||||
validity_date=invoice.validity_date,
|
||||
)
|
||||
|
||||
|
||||
def _publish_invoice_value(invoice: LightningInvoice, key: str, value: Any) -> None:
|
||||
"""Update a caller view without marking a mapped object dirty."""
|
||||
try:
|
||||
set_committed_value(invoice, key, value)
|
||||
except AttributeError:
|
||||
setattr(invoice, key, value)
|
||||
|
||||
|
||||
class InvoiceCreateRequest(BaseModel):
|
||||
amount_sats: int = Field(gt=0, le=1_000_000, description="Amount in satoshis")
|
||||
@@ -64,12 +115,80 @@ class InvoiceRecoverRequest(BaseModel):
|
||||
bolt11: str = Field(description="BOLT11 invoice string")
|
||||
|
||||
|
||||
def _trusted_mint_candidates() -> list[str]:
|
||||
return [
|
||||
mint
|
||||
for mint in dict.fromkeys([settings.primary_mint, *settings.cashu_mints])
|
||||
if mint
|
||||
]
|
||||
|
||||
|
||||
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.
|
||||
|
||||
Guards against amount_sats <= 0: the cashu library's PostMintQuoteRequest
|
||||
enforces ``amount > 0`` (Pydantic Field(gt=0)), so passing 0 raises a
|
||||
cryptic validation error deep in the stack. Fail fast with context.
|
||||
"""
|
||||
if amount_sats <= 0:
|
||||
raise ValueError(
|
||||
f"generate_lightning_invoice: amount_sats must be > 0, got {amount_sats}."
|
||||
)
|
||||
tried: list[str] = []
|
||||
configured = allowed_mints or [settings.primary_mint, *settings.cashu_mints]
|
||||
candidates = list(dict.fromkeys(configured))
|
||||
for mint_url in candidates:
|
||||
cooldown = _mint_cooldown_remaining(mint_url)
|
||||
if cooldown > 0:
|
||||
tried.append(f"{mint_url}: cooling down")
|
||||
logger.info(
|
||||
"Skipping rate-limited mint",
|
||||
extra={
|
||||
"mint_url": mint_url,
|
||||
"cooldown_seconds": round(cooldown, 2),
|
||||
"op_name": "request_mint_invoice",
|
||||
},
|
||||
)
|
||||
continue
|
||||
try:
|
||||
wallet = await get_wallet(mint_url, "sat", retry_on_rate_limit=False)
|
||||
quote = await _mint_operation(
|
||||
lambda: wallet.request_mint(amount_sats),
|
||||
op_name="request_mint_invoice",
|
||||
mint_url=mint_url,
|
||||
retry_on_rate_limit=False,
|
||||
)
|
||||
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 +202,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 +213,23 @@ 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
|
||||
allowed_mints = None
|
||||
if request.purpose == "topup":
|
||||
assert topup_api_key is not None
|
||||
# A key's liabilities are attributed to a single refund mint. Keep
|
||||
# top-up collateral on that same mint so balances and payouts cannot
|
||||
# misclassify funds held by another mint as owner profit.
|
||||
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 +244,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,
|
||||
@@ -222,82 +352,260 @@ async def recover_invoice(
|
||||
async def check_invoice_payment(
|
||||
invoice: LightningInvoice, session: AsyncSession
|
||||
) -> None:
|
||||
try:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
|
||||
mint_status = await wallet.get_mint_quote(invoice.payment_hash)
|
||||
|
||||
if mint_status.paid:
|
||||
invoice.status = "paid"
|
||||
invoice.paid_at = int(time.time())
|
||||
|
||||
if invoice.purpose == "create":
|
||||
api_key = await create_api_key_from_invoice(invoice, session)
|
||||
invoice.api_key_hash = api_key.hashed_key
|
||||
elif invoice.purpose == "topup" and invoice.api_key_hash:
|
||||
await topup_api_key_from_invoice(invoice, session)
|
||||
|
||||
lock = _invoice_settlement_locks.setdefault(invoice.id, asyncio.Lock())
|
||||
async with lock, wallet_operation_guard():
|
||||
minted = False
|
||||
try:
|
||||
# Snapshot the row and end the caller's read transaction before any
|
||||
# potentially slow mint I/O. All final DB mutations use owned,
|
||||
# short-lived sessions below.
|
||||
await session.refresh(invoice)
|
||||
if invoice.status != "pending":
|
||||
await session.commit()
|
||||
return
|
||||
settlement = _InvoiceSettlement.from_invoice(invoice)
|
||||
await session.commit()
|
||||
|
||||
mint_url = settlement.mint_url or settings.primary_mint
|
||||
wallet = await get_wallet(mint_url, "sat")
|
||||
mint_status = await _mint_operation(
|
||||
lambda: wallet.get_mint_quote(settlement.payment_hash),
|
||||
op_name="get_mint_quote",
|
||||
mint_url=mint_url,
|
||||
)
|
||||
if not mint_status.paid:
|
||||
return
|
||||
|
||||
# Reject a paid top-up whose target was pruned before redeeming its
|
||||
# single-use quote. The validation session is closed before mint I/O.
|
||||
if settlement.purpose == "topup":
|
||||
if not settlement.api_key_hash:
|
||||
raise ValueError("No API key associated with topup invoice")
|
||||
async with create_session() as validation_session:
|
||||
target = await validation_session.get(
|
||||
ApiKey, settlement.api_key_hash
|
||||
)
|
||||
if target is None:
|
||||
terminal = await validation_session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(
|
||||
col(LightningInvoice.id) == settlement.id,
|
||||
col(LightningInvoice.status) == "pending",
|
||||
)
|
||||
.values(status="reconciliation_required")
|
||||
)
|
||||
await validation_session.commit()
|
||||
if terminal.rowcount == 1:
|
||||
_publish_invoice_value(
|
||||
invoice, "status", "reconciliation_required"
|
||||
)
|
||||
else:
|
||||
await _reload_invoice_view(invoice, session)
|
||||
logger.critical(
|
||||
"Paid topup invoice target API key was not found; reconciliation required",
|
||||
extra={"invoice_id": settlement.id},
|
||||
)
|
||||
return
|
||||
|
||||
# Quote-linked proof verification makes an ambiguous mint response
|
||||
# retryable without crediting unrelated wallet balance growth.
|
||||
await _mint_invoice_quote(wallet, settlement)
|
||||
minted = True
|
||||
|
||||
paid_at = int(time.time())
|
||||
async with create_session() as finalization_session:
|
||||
settled, api_key_hash = await _finalize_invoice_settlement(
|
||||
settlement, finalization_session, paid_at
|
||||
)
|
||||
if not settled:
|
||||
await _reload_invoice_view(invoice, session)
|
||||
return
|
||||
|
||||
_publish_invoice_value(invoice, "status", "paid")
|
||||
_publish_invoice_value(invoice, "paid_at", paid_at)
|
||||
_publish_invoice_value(invoice, "api_key_hash", api_key_hash)
|
||||
logger.info(
|
||||
"Lightning invoice paid",
|
||||
extra={
|
||||
"invoice_id": invoice.id,
|
||||
"amount_sats": invoice.amount_sats,
|
||||
"purpose": invoice.purpose,
|
||||
"api_key_hash": invoice.api_key_hash[:8] + "..."
|
||||
if invoice.api_key_hash
|
||||
"invoice_id": settlement.id,
|
||||
"amount_sats": settlement.amount_sats,
|
||||
"purpose": settlement.purpose,
|
||||
"api_key_hash": api_key_hash[:8] + "..."
|
||||
if api_key_hash
|
||||
else None,
|
||||
},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to check invoice payment: {e}")
|
||||
except BaseException as error:
|
||||
# Never roll back the caller-owned session: doing so expires invoice
|
||||
# and sibling ORM objects. Owned sessions roll themselves back.
|
||||
if minted:
|
||||
logger.critical(
|
||||
"Invoice mint succeeded but DB finalization failed; reconciliation required",
|
||||
extra={"invoice_id": invoice.id, "purpose": invoice.purpose},
|
||||
)
|
||||
try:
|
||||
await _reload_invoice_view(invoice, session)
|
||||
except Exception:
|
||||
pass
|
||||
if not isinstance(error, Exception):
|
||||
raise
|
||||
logger.error(f"Failed to check invoice payment: {error}")
|
||||
|
||||
|
||||
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)
|
||||
def _is_outputs_already_signed(error: BaseException) -> bool:
|
||||
message = str(error)
|
||||
return bool(
|
||||
re.search(
|
||||
r"\boutputs?\s+(?:have\s+)?already\s+(?:been\s+)?signed(?:\s+before)?\b",
|
||||
message,
|
||||
re.IGNORECASE,
|
||||
)
|
||||
and re.search(r"\bcode\s*:\s*11003\b", message, re.IGNORECASE)
|
||||
)
|
||||
|
||||
|
||||
def _invoice_quote_proof_amount(wallet: Any, quote_id: str) -> int:
|
||||
"""Return spendable wallet value minted by one Lightning quote."""
|
||||
return sum(
|
||||
proof.amount
|
||||
for proof in wallet.proofs
|
||||
if proof.mint_id == quote_id and not proof.reserved
|
||||
)
|
||||
|
||||
|
||||
async def _mint_invoice_quote(
|
||||
wallet: Any, invoice: LightningInvoice | _InvoiceSettlement
|
||||
) -> None:
|
||||
"""Mint a paid quote, proving quote-linked outputs before DB credit."""
|
||||
mint_url = invoice.mint_url or settings.primary_mint
|
||||
await wallet.load_proofs(reload=True)
|
||||
if _invoice_quote_proof_amount(wallet, invoice.payment_hash) >= invoice.amount_sats:
|
||||
return
|
||||
|
||||
try:
|
||||
await _mint_operation(
|
||||
lambda: wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash),
|
||||
op_name=f"invoice_mint_{invoice.purpose}",
|
||||
mint_url=mint_url,
|
||||
retry_timeouts=False,
|
||||
)
|
||||
except Exception as error:
|
||||
if not _is_outputs_already_signed(error):
|
||||
raise
|
||||
|
||||
for keyset_id in wallet.keysets:
|
||||
await wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25)
|
||||
await wallet.load_proofs(reload=True)
|
||||
recovered = _invoice_quote_proof_amount(wallet, invoice.payment_hash)
|
||||
if recovered < invoice.amount_sats:
|
||||
raise RuntimeError(
|
||||
"Invoice outputs were already signed but quote-linked recovery returned "
|
||||
f"{recovered} sats; expected at least {invoice.amount_sats}"
|
||||
) from error
|
||||
else:
|
||||
await wallet.load_proofs(reload=True)
|
||||
minted_amount = _invoice_quote_proof_amount(wallet, invoice.payment_hash)
|
||||
if minted_amount < invoice.amount_sats:
|
||||
raise RuntimeError(
|
||||
"Invoice mint succeeded but quote-linked proofs total "
|
||||
f"{minted_amount} sats; expected at least {invoice.amount_sats}"
|
||||
)
|
||||
|
||||
|
||||
def _invoice_api_key_hash(invoice: LightningInvoice | _InvoiceSettlement) -> str:
|
||||
dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}"
|
||||
hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest()
|
||||
return hashlib.sha256(dummy_token.encode()).hexdigest()
|
||||
|
||||
|
||||
async def _create_api_key_record(
|
||||
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
|
||||
) -> ApiKey:
|
||||
mint_url = invoice.mint_url or settings.primary_mint
|
||||
api_key = ApiKey(
|
||||
hashed_key=hashed_key,
|
||||
balance=invoice.amount_sats * 1000, # Convert to msats
|
||||
hashed_key=_invoice_api_key_hash(invoice),
|
||||
balance=invoice.amount_sats * 1000,
|
||||
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,
|
||||
)
|
||||
|
||||
session.add(api_key)
|
||||
await session.flush()
|
||||
|
||||
return api_key
|
||||
|
||||
|
||||
async def topup_api_key_from_invoice(
|
||||
invoice: LightningInvoice, session: AsyncSession
|
||||
async def _topup_api_key_record(
|
||||
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
|
||||
) -> None:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash)
|
||||
|
||||
if not invoice.api_key_hash:
|
||||
raise ValueError("No API key associated with topup invoice")
|
||||
|
||||
api_key = await session.get(ApiKey, invoice.api_key_hash)
|
||||
if not api_key:
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == invoice.api_key_hash)
|
||||
.values(balance=col(ApiKey.balance) + invoice.amount_sats * 1000)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
if result.rowcount != 1:
|
||||
raise ValueError("Associated API key not found")
|
||||
|
||||
api_key.balance += invoice.amount_sats * 1000 # Convert to msats
|
||||
await session.flush()
|
||||
|
||||
async def _finalize_invoice_settlement(
|
||||
invoice: _InvoiceSettlement, session: AsyncSession, paid_at: int
|
||||
) -> tuple[bool, str | None]:
|
||||
"""Atomically fence and apply one invoice credit in the provided owned session."""
|
||||
api_key_hash = (
|
||||
_invoice_api_key_hash(invoice)
|
||||
if invoice.purpose == "create"
|
||||
else invoice.api_key_hash
|
||||
)
|
||||
claim = await session.exec( # type: ignore[call-overload]
|
||||
update(LightningInvoice)
|
||||
.where(col(LightningInvoice.id) == invoice.id)
|
||||
.where(col(LightningInvoice.status) == "pending")
|
||||
.values(status="paid", paid_at=paid_at, api_key_hash=api_key_hash)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
if claim.rowcount != 1:
|
||||
await session.rollback()
|
||||
return False, None
|
||||
|
||||
if invoice.purpose == "create":
|
||||
await _create_api_key_record(invoice, session)
|
||||
elif invoice.purpose == "topup":
|
||||
await _topup_api_key_record(invoice, session)
|
||||
else:
|
||||
raise ValueError(f"Unsupported invoice purpose: {invoice.purpose}")
|
||||
await session.commit()
|
||||
return True, api_key_hash
|
||||
|
||||
|
||||
INVOICE_WATCH_INTERVAL_SECONDS = 5
|
||||
async def _reload_invoice_view(
|
||||
invoice: LightningInvoice, _caller_session: AsyncSession
|
||||
) -> None:
|
||||
"""Publish committed invoice state without touching the caller transaction."""
|
||||
async with create_session() as reload_session:
|
||||
stored = await reload_session.get(LightningInvoice, invoice.id)
|
||||
if stored is None:
|
||||
return
|
||||
status = stored.status
|
||||
paid_at = stored.paid_at
|
||||
api_key_hash = stored.api_key_hash
|
||||
await reload_session.commit()
|
||||
_publish_invoice_value(invoice, "status", status)
|
||||
_publish_invoice_value(invoice, "paid_at", paid_at)
|
||||
_publish_invoice_value(invoice, "api_key_hash", api_key_hash)
|
||||
|
||||
|
||||
async def _credit_topup_record(
|
||||
invoice: LightningInvoice | _InvoiceSettlement, session: AsyncSession
|
||||
) -> None:
|
||||
await _topup_api_key_record(invoice, session)
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
|
||||
|
||||
@@ -18,6 +18,17 @@ from ..wallet import deserialize_token_from_string
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Interim policy: when Routstr must move value to another trusted mint, the
|
||||
# cross-mint Lightning round trip can consume fees that are not visible to the
|
||||
# client. Reserve 5% headroom until the fee-payer policy is made explicit.
|
||||
_MINT_FEE_ALLOWANCE = 0.05
|
||||
|
||||
|
||||
def apply_mint_fee_allowance(cost_msat: int) -> int:
|
||||
"""Reserve headroom for possible trusted-mint fallback fees."""
|
||||
adjusted = math.ceil(cost_msat * (1 - _MINT_FEE_ALLOWANCE))
|
||||
return max(settings.min_request_msat, adjusted)
|
||||
|
||||
|
||||
def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> None:
|
||||
if x_cashu := headers.get("x-cashu", None):
|
||||
|
||||
@@ -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
|
||||
|
||||
+52
-12
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
@@ -25,8 +26,8 @@ from .core.db import (
|
||||
)
|
||||
from .core.exceptions import UpstreamError
|
||||
from .core.not_found import build_not_found_response
|
||||
from .core.settings import settings
|
||||
from .payment.helpers import (
|
||||
apply_mint_fee_allowance,
|
||||
calculate_discounted_max_cost,
|
||||
check_token_balance,
|
||||
create_error_response,
|
||||
@@ -49,6 +50,13 @@ _provider_map: dict[
|
||||
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
|
||||
|
||||
|
||||
async def _finish_read_transaction(session: AsyncSession) -> None:
|
||||
"""Release a read transaction without assuming a particular session mock."""
|
||||
commit_result = session.commit()
|
||||
if inspect.isawaitable(commit_result):
|
||||
await commit_result
|
||||
|
||||
|
||||
async def initialize_upstreams() -> None:
|
||||
"""Initialize upstream providers from database during application startup."""
|
||||
global _upstreams
|
||||
@@ -220,6 +228,20 @@ _API_PATH_PREFIXES = (
|
||||
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||
async def proxy(
|
||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||
) -> Response | StreamingResponse:
|
||||
"""Run proxy setup in a short request session, never across response streaming."""
|
||||
try:
|
||||
return await _proxy(request, path, session)
|
||||
finally:
|
||||
# FastAPI yield dependencies normally close after the response body is
|
||||
# sent. Close explicitly so a long stream cannot retain DB resources.
|
||||
close_result = session.close()
|
||||
if inspect.isawaitable(close_result):
|
||||
await close_result
|
||||
|
||||
|
||||
async def _proxy(
|
||||
request: Request, path: str, session: AsyncSession
|
||||
) -> Response | StreamingResponse:
|
||||
# GET requests must hit a known API prefix; otherwise return a 404 (HTML
|
||||
# for browsers, JSON for API clients). POST requests are always forwarded
|
||||
@@ -332,8 +354,7 @@ async def proxy(
|
||||
max_cost_for_model = await calculate_discounted_max_cost(
|
||||
_max_cost_for_model, request_body_dict, model_obj=model_obj
|
||||
)
|
||||
# Ensure max_cost_for_model is at least the minimum allowed request cost
|
||||
max_cost_for_model = max(max_cost_for_model, settings.min_request_msat)
|
||||
max_cost_for_model = apply_mint_fee_allowance(max_cost_for_model)
|
||||
|
||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||
|
||||
@@ -450,6 +471,9 @@ async def proxy(
|
||||
if is_ehbp or request_body_dict:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
# Snapshot validation performs SELECTs after pay_for_request commits.
|
||||
# End that read transaction before waiting on upstream response headers.
|
||||
await _finish_read_transaction(session)
|
||||
|
||||
# Tracks request params already removed in response to upstream rejections,
|
||||
# shared across providers so a stripped param stays stripped on failover and
|
||||
@@ -469,7 +493,9 @@ async def proxy(
|
||||
candidate_max = await calculate_discounted_max_cost(
|
||||
candidate_max, request_body_dict, model_obj=model_obj
|
||||
)
|
||||
candidate_max = max(candidate_max, settings.min_request_msat)
|
||||
# Apply the same interim 5% trusted-mint fee headroom used for the
|
||||
# first candidate; failover must not silently change admission.
|
||||
candidate_max = apply_mint_fee_allowance(candidate_max)
|
||||
if candidate_max > max_cost_for_model:
|
||||
await revert_pay_for_request(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
@@ -481,8 +507,10 @@ async def proxy(
|
||||
raise
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
await _finish_read_transaction(session)
|
||||
continue
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
await _finish_read_transaction(session)
|
||||
max_cost_for_model = candidate_max
|
||||
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
@@ -767,17 +795,29 @@ async def get_bearer_token_key(
|
||||
},
|
||||
)
|
||||
return key
|
||||
except Exception as e:
|
||||
key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key
|
||||
logger.error(
|
||||
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}",
|
||||
except HTTPException as error:
|
||||
detail: dict[str, Any] = error.detail if isinstance(error.detail, dict) else {}
|
||||
raw_error = detail.get("error")
|
||||
error_info = raw_error if isinstance(raw_error, dict) else {}
|
||||
logger.warning(
|
||||
"Bearer token rejected",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"status_code": error.status_code,
|
||||
"error_code": error_info.get("code"),
|
||||
"path": path,
|
||||
"model_id": model_id,
|
||||
"min_cost_msat": min_cost,
|
||||
"bearer_key_preview": key_preview,
|
||||
"required_msat": min_cost,
|
||||
},
|
||||
)
|
||||
raise
|
||||
except Exception as error:
|
||||
logger.exception(
|
||||
"Bearer token validation failed",
|
||||
extra={
|
||||
"error_type": type(error).__name__,
|
||||
"path": path,
|
||||
"model_id": model_id,
|
||||
"required_msat": min_cost,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -12,7 +12,7 @@ from ..core.db import (
|
||||
from ..core.db import (
|
||||
store_cashu_transaction_with_retry as store_cashu_transaction,
|
||||
)
|
||||
from ..wallet import send_token
|
||||
from ..wallet import release_token_reservation, send_token, token_mint_url
|
||||
from .routstr import RoutstrUpstreamProvider
|
||||
|
||||
logger = get_logger(__name__)
|
||||
@@ -144,12 +144,13 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
)
|
||||
return
|
||||
|
||||
actual_mint_url = token_mint_url(token, mint_url)
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit="sat",
|
||||
mint_url=mint_url,
|
||||
mint_url=actual_mint_url,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
@@ -157,8 +158,24 @@ async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
except Exception:
|
||||
logger.critical(
|
||||
"Aborting auto top-up because its cashu token could not be persisted",
|
||||
extra={"provider_id": row.id, "mint_url": mint_url},
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
try:
|
||||
await release_token_reservation(token)
|
||||
except Exception as error:
|
||||
logger.critical(
|
||||
"Failed to release untracked auto-topup token",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"mint_url": actual_mint_url,
|
||||
"error": str(error),
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Auto-topup token was released after persistence failed",
|
||||
extra={"provider_id": row.id, "mint_url": actual_mint_url},
|
||||
)
|
||||
return
|
||||
|
||||
result = await provider.topup(token)
|
||||
|
||||
@@ -53,6 +53,7 @@ from ..wallet import (
|
||||
classify_redemption_error,
|
||||
recieve_token,
|
||||
send_token,
|
||||
token_mint_url,
|
||||
)
|
||||
from . import messages_dispatch
|
||||
from .cache_breakpoints import (
|
||||
@@ -3520,7 +3521,7 @@ class BaseUpstreamProvider:
|
||||
token=refund_token,
|
||||
amount=amount,
|
||||
unit=unit,
|
||||
mint_url=mint,
|
||||
mint_url=token_mint_url(refund_token, mint),
|
||||
typ="out",
|
||||
request_id=request_id,
|
||||
)
|
||||
@@ -3873,7 +3874,7 @@ class BaseUpstreamProvider:
|
||||
token=refund_token,
|
||||
amount=emergency_refund,
|
||||
unit=unit,
|
||||
mint_url=mint,
|
||||
mint_url=token_mint_url(refund_token, mint),
|
||||
typ="out",
|
||||
request_id=request_id,
|
||||
)
|
||||
@@ -4843,7 +4844,7 @@ class BaseUpstreamProvider:
|
||||
token=refund_token,
|
||||
amount=emergency_refund,
|
||||
unit=unit,
|
||||
mint_url=mint,
|
||||
mint_url=token_mint_url(refund_token, mint),
|
||||
typ="out",
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
+1421
-276
File diff suppressed because it is too large
Load Diff
@@ -207,8 +207,30 @@ async def test_pay_for_request_succeeds_when_balance_equals_cost(
|
||||
assert key.balance == model_cost # balance unchanged, only reserved goes up
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_five_percent_mint_fallback_headroom_is_admitted_and_reserved(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
from routstr.auth import pay_for_request, validate_bearer_key
|
||||
from routstr.payment.helpers import apply_mint_fee_allowance
|
||||
|
||||
key = _key(balance=95_000)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
admission_cost = apply_mint_fee_allowance(100_000)
|
||||
validated = await validate_bearer_key(
|
||||
f"sk-{key.hashed_key}", integration_session, min_cost=admission_cost
|
||||
)
|
||||
await pay_for_request(validated, admission_cost, integration_session)
|
||||
|
||||
await integration_session.refresh(key)
|
||||
assert admission_cost == 95_000
|
||||
assert key.reserved_balance == 95_000
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 6 — HTTP layer returns 402 JSON with the right shape
|
||||
# HTTP layer returns 402 JSON with the right shape
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -266,8 +288,8 @@ async def test_http_402_response_shape_on_insufficient_balance(
|
||||
error = body["detail"]["error"]
|
||||
assert error["code"] == "insufficient_balance"
|
||||
assert error["type"] == "insufficient_quota"
|
||||
assert str(model_cost) in error["message"]
|
||||
assert str(user_balance) in error["message"]
|
||||
assert "591.744 sats (591744 msats) required" in error["message"]
|
||||
assert "20.32 sats (20320 msats) available" in error["message"]
|
||||
|
||||
# Balance must be completely untouched
|
||||
await integration_session.refresh(key)
|
||||
|
||||
@@ -3,20 +3,30 @@
|
||||
Covers two things:
|
||||
- The three constraint fields (balance_limit, balance_limit_reset, validity_date)
|
||||
are persisted on LightningInvoice and survive a DB round-trip.
|
||||
- create_api_key_from_invoice propagates those fields to the created ApiKey,
|
||||
so the constraints are actually enforced when the key is used.
|
||||
- The production-path API-key record helper propagates those fields to the
|
||||
created ApiKey, so the constraints are actually enforced when the key is used.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
from sqlalchemy import inspect
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey, LightningInvoice
|
||||
from routstr.lightning import create_api_key_from_invoice
|
||||
from routstr.lightning import _create_api_key_record
|
||||
|
||||
|
||||
def _configure_quote_proof_wallet(wallet: MagicMock) -> None:
|
||||
wallet.proofs = []
|
||||
wallet.keysets = {}
|
||||
wallet.load_proofs = AsyncMock()
|
||||
|
||||
|
||||
def _make_invoice(**kwargs: object) -> LightningInvoice:
|
||||
@@ -39,7 +49,15 @@ def _make_invoice(**kwargs: object) -> LightningInvoice:
|
||||
def mock_wallet_mint() -> object:
|
||||
with patch("routstr.lightning.get_wallet") as mock_get_wallet:
|
||||
wallet = AsyncMock()
|
||||
wallet.mint = AsyncMock(return_value=[])
|
||||
wallet.proofs = []
|
||||
wallet.load_proofs = AsyncMock()
|
||||
|
||||
async def mint(amount: int, quote_id: str) -> list[Proof]:
|
||||
proofs = [Proof(amount=amount, mint_id=quote_id)]
|
||||
wallet.proofs.extend(proofs)
|
||||
return proofs
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=mint)
|
||||
mock_get_wallet.return_value = wallet
|
||||
yield mock_get_wallet
|
||||
|
||||
@@ -48,6 +66,7 @@ def mock_wallet_mint() -> object:
|
||||
# Persistence
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_persists_balance_limit(
|
||||
integration_session: AsyncSession,
|
||||
@@ -92,6 +111,7 @@ async def test_invoice_persists_validity_date(
|
||||
# Propagation to ApiKey
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_created_key_receives_balance_limit(
|
||||
integration_session: AsyncSession,
|
||||
@@ -100,7 +120,7 @@ async def test_created_key_receives_balance_limit(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -116,7 +136,7 @@ async def test_created_key_receives_balance_limit_reset(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -133,7 +153,7 @@ async def test_created_key_receives_validity_date(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -141,6 +161,254 @@ async def test_created_key_receives_validity_date(
|
||||
assert stored_key.validity_date == expiry
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_payment_check_releases_connection_during_mint_quote(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_slow_quote", status="pending", paid_at=None)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
|
||||
async def quote_status(*args: object, **kwargs: object) -> MagicMock:
|
||||
assert integration_engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return MagicMock(paid=False)
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_mint_quote = AsyncMock(side_effect=quote_status)
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_payment_checks_mint_and_credit_invoice_once(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_concurrent", status="pending", paid_at=None)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
|
||||
mint_calls = 0
|
||||
|
||||
async def single_use_mint(*args: object, **kwargs: object) -> list[object]:
|
||||
# Real mints enforce single-use quotes: the second concurrent minter
|
||||
# gets rejected at the mint, mirroring cashu quote semantics.
|
||||
nonlocal mint_calls
|
||||
mint_calls += 1
|
||||
call_number = mint_calls
|
||||
await asyncio.sleep(0.05)
|
||||
if call_number > 1:
|
||||
raise Exception("quote already issued")
|
||||
proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash)
|
||||
wallet.proofs.append(proof)
|
||||
return [proof]
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=single_use_mint)
|
||||
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as first,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as second,
|
||||
):
|
||||
first_invoice = await first.get(LightningInvoice, invoice.id)
|
||||
second_invoice = await second.get(LightningInvoice, invoice.id)
|
||||
assert first_invoice is not None
|
||||
assert second_invoice is not None
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await asyncio.gather(
|
||||
check_invoice_payment(first_invoice, first),
|
||||
check_invoice_payment(second_invoice, second),
|
||||
)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_invoice = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored_invoice is not None
|
||||
assert stored_invoice.status == "paid"
|
||||
assert stored_invoice.api_key_hash is not None
|
||||
stored_key = await verify.get(ApiKey, stored_invoice.api_key_hash)
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == invoice.amount_sats * 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_mint_keeps_invoice_pending_for_retry(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_mint_failure", status="pending", paid_at=None)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
wallet.mint = AsyncMock(side_effect=TimeoutError("mint unavailable"))
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "pending"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unpaid_topup_does_not_query_target_key(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(
|
||||
id="inv_unpaid_topup",
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
purpose="topup",
|
||||
api_key_hash="target-key",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=False))
|
||||
create_session = MagicMock(side_effect=RuntimeError("target lookup should not run"))
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.create_session", create_session),
|
||||
):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
wallet.get_mint_quote.assert_awaited_once_with(invoice.payment_hash)
|
||||
create_session.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_topup_target_is_rejected_before_mint(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(
|
||||
id="inv_missing_topup_target",
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
purpose="topup",
|
||||
api_key_hash="pruned-key",
|
||||
expires_at=int(time.time()) - 1,
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(invoice)
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
wallet.mint = AsyncMock()
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.logger.critical") as critical,
|
||||
):
|
||||
from routstr.lightning import get_invoice_status
|
||||
|
||||
response = await get_invoice_status(invoice.id, session)
|
||||
|
||||
assert response.status == "reconciliation_required"
|
||||
assert stored.status == "reconciliation_required"
|
||||
assert stored not in session.dirty
|
||||
critical.assert_called_once()
|
||||
|
||||
wallet.mint.assert_not_awaited()
|
||||
async with AsyncSession(integration_engine) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "reconciliation_required"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_mint_db_failure_keeps_invoice_pending_for_reconciliation(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
invoice = _make_invoice(id="inv_finalize_failure", status="pending", paid_at=None)
|
||||
sibling = _make_invoice(
|
||||
id="inv_finalize_failure_sibling",
|
||||
bolt11="lnbc1000n1sibling",
|
||||
payment_hash="cafebabe" * 8,
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add_all([invoice, sibling])
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
|
||||
async def successful_mint(*args: object, **kwargs: object) -> list[Proof]:
|
||||
proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash)
|
||||
wallet.proofs.append(proof)
|
||||
return [proof]
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=successful_mint)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as session:
|
||||
stored = await session.get(LightningInvoice, invoice.id)
|
||||
stored_sibling = await session.get(LightningInvoice, sibling.id)
|
||||
assert stored is not None
|
||||
assert stored_sibling is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.lightning._create_api_key_record",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
),
|
||||
):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await check_invoice_payment(stored, session)
|
||||
|
||||
stored_state = inspect(stored)
|
||||
sibling_state = inspect(stored_sibling)
|
||||
assert stored_state is not None
|
||||
assert sibling_state is not None
|
||||
assert stored_state.expired is False
|
||||
assert sibling_state.expired is False
|
||||
assert stored.status == "pending"
|
||||
assert stored_sibling.id == sibling.id
|
||||
|
||||
assert wallet.mint.await_count == 1
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
assert stored.status == "pending"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_created_key_without_constraints_has_none_fields(
|
||||
integration_session: AsyncSession,
|
||||
@@ -149,7 +417,7 @@ async def test_created_key_without_constraints_has_none_fields(
|
||||
integration_session.add(invoice)
|
||||
await integration_session.flush()
|
||||
|
||||
api_key = await create_api_key_from_invoice(invoice, integration_session)
|
||||
api_key = await _create_api_key_record(invoice, integration_session)
|
||||
await integration_session.commit()
|
||||
|
||||
stored_key = await integration_session.get(ApiKey, api_key.hashed_key)
|
||||
@@ -157,3 +425,87 @@ async def test_created_key_without_constraints_has_none_fields(
|
||||
assert stored_key.balance_limit is None
|
||||
assert stored_key.balance_limit_reset is None
|
||||
assert stored_key.validity_date is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_guard_credits_once_when_both_mints_succeed(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Even if the mint fails to enforce single-use quotes and both racers
|
||||
mint successfully, the conditional status update must credit exactly once."""
|
||||
key = ApiKey(hashed_key="race-key", balance=1_000)
|
||||
invoice = _make_invoice(
|
||||
id="inv_db_guard",
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
purpose="topup",
|
||||
api_key_hash="race-key",
|
||||
)
|
||||
sibling = _make_invoice(
|
||||
id="inv_db_guard_sibling",
|
||||
bolt11="lnbc1000n1race-sibling",
|
||||
payment_hash="01234567" * 8,
|
||||
status="pending",
|
||||
paid_at=None,
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add_all([key, invoice, sibling])
|
||||
await setup.commit()
|
||||
|
||||
wallet = MagicMock()
|
||||
_configure_quote_proof_wallet(wallet)
|
||||
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
|
||||
|
||||
async def always_succeeding_mint(*args: object, **kwargs: object) -> list[Proof]:
|
||||
await asyncio.sleep(0.05)
|
||||
proof = Proof(amount=invoice.amount_sats, mint_id=invoice.payment_hash)
|
||||
wallet.proofs.append(proof)
|
||||
return [proof]
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=always_succeeding_mint)
|
||||
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as first,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as second,
|
||||
):
|
||||
first_invoice = await first.get(LightningInvoice, invoice.id)
|
||||
first_sibling = await first.get(LightningInvoice, sibling.id)
|
||||
second_invoice = await second.get(LightningInvoice, invoice.id)
|
||||
assert first_invoice is not None
|
||||
assert first_sibling is not None
|
||||
assert second_invoice is not None
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
from routstr.lightning import check_invoice_payment
|
||||
|
||||
await asyncio.gather(
|
||||
check_invoice_payment(first_invoice, first),
|
||||
check_invoice_payment(second_invoice, second),
|
||||
)
|
||||
|
||||
first_state = inspect(first_invoice)
|
||||
sibling_state = inspect(first_sibling)
|
||||
second_state = inspect(second_invoice)
|
||||
assert first_state is not None
|
||||
assert sibling_state is not None
|
||||
assert second_state is not None
|
||||
assert first_state.expired is False
|
||||
assert sibling_state.expired is False
|
||||
assert second_state.expired is False
|
||||
assert first_invoice.id == invoice.id
|
||||
assert first_sibling.id == sibling.id
|
||||
assert second_invoice.id == invoice.id
|
||||
assert first_invoice.status == "paid"
|
||||
assert second_invoice.status == "paid"
|
||||
assert first_invoice not in first.dirty
|
||||
assert second_invoice not in second.dirty
|
||||
|
||||
assert wallet.mint.await_count == 1
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_invoice = await verify.get(LightningInvoice, invoice.id)
|
||||
assert stored_invoice is not None
|
||||
assert stored_invoice.status == "paid"
|
||||
stored_key = await verify.get(ApiKey, "race-key")
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == 1_000 + invoice.amount_sats * 1000
|
||||
|
||||
@@ -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,8 @@ async def test_topup_with_authorization_header(
|
||||
body = resp.json()
|
||||
assert body["amount_sats"] == 500
|
||||
assert body["bolt11"].startswith("lnbc")
|
||||
allowed_mints = patch_invoice_generation.call_args.kwargs["allowed_mints"]
|
||||
assert allowed_mints == ["http://localhost:3338"]
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
|
||||
@@ -0,0 +1,273 @@
|
||||
import asyncio
|
||||
import time
|
||||
import uuid
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel import col, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey, LightningInvoice
|
||||
from routstr.lightning import (
|
||||
_finalize_invoice_settlement,
|
||||
_InvoiceSettlement,
|
||||
check_invoice_payment,
|
||||
)
|
||||
|
||||
|
||||
def _lightning_invoice(**overrides: object) -> LightningInvoice:
|
||||
suffix = uuid.uuid4().hex
|
||||
values = {
|
||||
"id": f"invoice-{suffix}",
|
||||
"bolt11": f"lnbc-{suffix}",
|
||||
"amount_sats": 100,
|
||||
"description": "settlement test",
|
||||
"payment_hash": f"quote-{suffix}",
|
||||
"status": "pending",
|
||||
"purpose": "create",
|
||||
"mint_url": "http://mint:3338",
|
||||
"expires_at": int(time.time()) + 3600,
|
||||
}
|
||||
values.update(overrides)
|
||||
return LightningInvoice(**values) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_read_transaction_closes_before_external_mint_io(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
invoice = _lightning_invoice()
|
||||
integration_session.add(invoice)
|
||||
await integration_session.commit()
|
||||
stored = await integration_session.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
|
||||
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=False)))
|
||||
|
||||
async def get_wallet_without_open_db_transaction(
|
||||
*args: object, **kwargs: object
|
||||
) -> Mock:
|
||||
assert not integration_session.in_transaction()
|
||||
return wallet
|
||||
|
||||
with patch(
|
||||
"routstr.lightning.get_wallet", side_effect=get_wallet_without_open_db_transaction
|
||||
):
|
||||
await check_invoice_payment(stored, integration_session)
|
||||
|
||||
assert not integration_session.in_transaction()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_separate_sessions_cas_topup_credit_exactly_once(
|
||||
integration_engine: AsyncEngine,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup",
|
||||
api_key_hash=key_hash,
|
||||
amount_sats=100,
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
snapshot_a = _InvoiceSettlement.from_invoice(invoice)
|
||||
snapshot_b = _InvoiceSettlement.from_invoice(invoice)
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as session_a,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as session_b,
|
||||
):
|
||||
results = await asyncio.gather(
|
||||
_finalize_invoice_settlement(snapshot_a, session_a, 1_700_000_000),
|
||||
_finalize_invoice_settlement(snapshot_b, session_b, 1_700_000_001),
|
||||
)
|
||||
|
||||
assert sorted(settled for settled, _ in results) == [False, True]
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_invoice = await verify.get(LightningInvoice, invoice.id)
|
||||
stored_key = await verify.get(ApiKey, key_hash)
|
||||
assert stored_invoice is not None
|
||||
assert stored_invoice.status == "paid"
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == 200_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_atomic_increment_preserves_concurrent_balance_mutation(
|
||||
integration_engine: AsyncEngine,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup", api_key_hash=key_hash, amount_sats=100
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
async def debit_balance(session: AsyncSession) -> None:
|
||||
result = await session.exec( # type: ignore[call-overload]
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == key_hash)
|
||||
.values(balance=col(ApiKey.balance) - 10_000)
|
||||
.execution_options(synchronize_session=False)
|
||||
)
|
||||
assert result.rowcount == 1
|
||||
await session.commit()
|
||||
|
||||
snapshot = _InvoiceSettlement.from_invoice(invoice)
|
||||
async with (
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as settlement,
|
||||
AsyncSession(integration_engine, expire_on_commit=False) as debit,
|
||||
):
|
||||
settlement_result, _ = await asyncio.gather(
|
||||
_finalize_invoice_settlement(snapshot, settlement, 1_700_000_000),
|
||||
debit_balance(debit),
|
||||
)
|
||||
|
||||
assert settlement_result[0]
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
stored_key = await verify.get(ApiKey, key_hash)
|
||||
assert stored_key is not None
|
||||
assert stored_key.balance == 190_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_final_commit_rolls_back_claim_and_credit_for_retry(
|
||||
integration_engine: AsyncEngine,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup",
|
||||
api_key_hash=key_hash,
|
||||
amount_sats=100,
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
snapshot = _InvoiceSettlement.from_invoice(invoice)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as failed:
|
||||
with patch.object(
|
||||
failed, "commit", AsyncMock(side_effect=Exception("db unavailable"))
|
||||
):
|
||||
with pytest.raises(Exception, match="db unavailable"):
|
||||
await _finalize_invoice_settlement(snapshot, failed, 1_700_000_000)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
pending = await verify.get(LightningInvoice, invoice.id)
|
||||
unchanged = await verify.get(ApiKey, key_hash)
|
||||
assert pending is not None
|
||||
assert pending.status == "pending"
|
||||
assert unchanged is not None
|
||||
assert unchanged.balance == 100_000
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as retry:
|
||||
settled, _ = await _finalize_invoice_settlement(
|
||||
snapshot, retry, 1_700_000_001
|
||||
)
|
||||
assert settled
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
paid = await verify.get(LightningInvoice, invoice.id)
|
||||
credited = await verify.get(ApiKey, key_hash)
|
||||
assert paid is not None
|
||||
assert paid.status == "paid"
|
||||
assert credited is not None
|
||||
assert credited.balance == 200_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_invoice_payment_retries_after_mint_success_and_db_failure(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
key_hash = uuid.uuid4().hex
|
||||
invoice = _lightning_invoice(
|
||||
purpose="topup", api_key_hash=key_hash, amount_sats=100
|
||||
)
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=100_000,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://mint:3338",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as seed:
|
||||
seed.add(key)
|
||||
seed.add(invoice)
|
||||
await seed.commit()
|
||||
|
||||
wallet = Mock(
|
||||
proofs=[],
|
||||
keysets={"keyset-1": Mock()},
|
||||
load_proofs=AsyncMock(),
|
||||
get_mint_quote=AsyncMock(return_value=Mock(paid=True)),
|
||||
restore_tokens_for_keyset=AsyncMock(),
|
||||
)
|
||||
|
||||
async def mint(amount: int, quote_id: str) -> list[Proof]:
|
||||
proofs = [Proof(amount=amount, mint_id=quote_id)]
|
||||
wallet.proofs.extend(proofs)
|
||||
return proofs
|
||||
|
||||
wallet.mint = AsyncMock(side_effect=mint)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as failed:
|
||||
stored = await failed.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.lightning._finalize_invoice_settlement",
|
||||
AsyncMock(side_effect=Exception("db unavailable")),
|
||||
),
|
||||
):
|
||||
await check_invoice_payment(stored, failed)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
pending = await verify.get(LightningInvoice, invoice.id)
|
||||
unchanged = await verify.get(ApiKey, key_hash)
|
||||
assert pending is not None
|
||||
assert pending.status == "pending"
|
||||
assert unchanged is not None
|
||||
assert unchanged.balance == 100_000
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as retry:
|
||||
stored = await retry.get(LightningInvoice, invoice.id)
|
||||
assert stored is not None
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)):
|
||||
await check_invoice_payment(stored, retry)
|
||||
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as verify:
|
||||
paid = await verify.get(LightningInvoice, invoice.id)
|
||||
credited = await verify.get(ApiKey, key_hash)
|
||||
assert paid is not None
|
||||
assert paid.status == "paid"
|
||||
assert credited is not None
|
||||
assert credited.balance == 200_000
|
||||
|
||||
wallet.mint.assert_awaited_once_with(100, quote_id=invoice.payment_hash)
|
||||
wallet.restore_tokens_for_keyset.assert_not_awaited()
|
||||
@@ -0,0 +1,174 @@
|
||||
"""Money-safety regression coverage for automatic wallet payouts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core import db
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import credit_balance, periodic_payout
|
||||
|
||||
PRIMARY_MINT = "http://primary:3338"
|
||||
REFUND_MINT = "http://refund:3338"
|
||||
PAYOUT_INTERVAL = 987
|
||||
|
||||
|
||||
class _LoopBreak(Exception):
|
||||
"""Stop the otherwise-infinite payout loop after one cycle."""
|
||||
|
||||
|
||||
def _one_payout_cycle() -> Callable[[float], Coroutine[Any, Any, None]]:
|
||||
intervals_seen = 0
|
||||
|
||||
async def sleep(seconds: float) -> None:
|
||||
nonlocal intervals_seen
|
||||
if seconds == PAYOUT_INTERVAL:
|
||||
intervals_seen += 1
|
||||
if intervals_seen == 2:
|
||||
raise _LoopBreak()
|
||||
|
||||
return sleep
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cross_mint_liability_is_not_paid_as_owner_profit(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Refund preferences must not make primary-mint customer funds payable."""
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(
|
||||
ApiKey(
|
||||
hashed_key="cross-mint-key",
|
||||
balance=50_000,
|
||||
refund_mint_url=REFUND_MINT,
|
||||
refund_currency="sat",
|
||||
)
|
||||
)
|
||||
await setup.commit()
|
||||
|
||||
primary_proof = MagicMock(amount=50)
|
||||
raw_send = AsyncMock(return_value=50)
|
||||
|
||||
def proofs_for_mint(
|
||||
_wallet: object, mint_url: str, unit: str, **_kwargs: object
|
||||
) -> list[MagicMock]:
|
||||
if mint_url == PRIMARY_MINT and unit == "sat":
|
||||
return [primary_proof]
|
||||
return []
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", [REFUND_MINT]),
|
||||
patch.object(settings, "primary_mint", PRIMARY_MINT),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.test"),
|
||||
patch.object(settings, "payout_interval_seconds", PAYOUT_INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_payout_cycle()),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(side_effect=proofs_for_mint),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, _wallet: proofs),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
raw_send.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_payout_does_not_send_proofs_whose_liability_commit_is_in_flight(
|
||||
integration_engine: AsyncEngine,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Proof visibility before liability commit must not expose customer funds."""
|
||||
key = ApiKey(
|
||||
hashed_key="in-flight-topup-key",
|
||||
balance=0,
|
||||
refund_mint_url=PRIMARY_MINT,
|
||||
refund_currency="sat",
|
||||
)
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as setup:
|
||||
setup.add(key)
|
||||
await setup.commit()
|
||||
|
||||
proofs: list[MagicMock] = []
|
||||
proof_visible = asyncio.Event()
|
||||
finish_redemption = asyncio.Event()
|
||||
liability_read = asyncio.Event()
|
||||
|
||||
async def redeem_token(token: str) -> tuple[int, str, str]:
|
||||
proofs.append(MagicMock(amount=200))
|
||||
proof_visible.set()
|
||||
await finish_redemption.wait()
|
||||
return 200, "sat", PRIMARY_MINT
|
||||
|
||||
real_total_liability = db.total_user_liability
|
||||
|
||||
async def read_liability(_session: AsyncSession) -> int:
|
||||
async with db.create_session() as snapshot_session:
|
||||
value = await real_total_liability(snapshot_session)
|
||||
liability_read.set()
|
||||
return value
|
||||
|
||||
raw_send = AsyncMock(return_value=200)
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", PRIMARY_MINT),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.test"),
|
||||
patch.object(settings, "payout_interval_seconds", PAYOUT_INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_payout_cycle()),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=redeem_token)),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(side_effect=lambda *_args, **_kwargs: list(proofs)),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda visible, _wallet: visible),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.total_user_liability",
|
||||
AsyncMock(side_effect=read_liability),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
async with AsyncSession(integration_engine, expire_on_commit=False) as credit_session:
|
||||
stored_key = await credit_session.get(ApiKey, key.hashed_key)
|
||||
assert stored_key is not None
|
||||
credit_task = asyncio.create_task(
|
||||
credit_balance("cashu-token", stored_key, credit_session)
|
||||
)
|
||||
await asyncio.wait_for(proof_visible.wait(), timeout=2)
|
||||
|
||||
payout_task = asyncio.create_task(periodic_payout())
|
||||
try:
|
||||
await asyncio.wait_for(liability_read.wait(), timeout=0.1)
|
||||
liability_was_read_while_crediting = True
|
||||
except TimeoutError:
|
||||
liability_was_read_while_crediting = False
|
||||
|
||||
finish_redemption.set()
|
||||
await asyncio.wait_for(credit_task, timeout=2)
|
||||
|
||||
with pytest.raises(_LoopBreak):
|
||||
await asyncio.wait_for(payout_task, timeout=2)
|
||||
|
||||
assert liability_was_read_while_crediting is False
|
||||
raw_send.assert_not_awaited()
|
||||
@@ -0,0 +1,65 @@
|
||||
"""Integration coverage for proxy database-session lifetime."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Response
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticated_proxy_releases_db_connection_before_upstream_headers(
|
||||
integration_engine: AsyncEngine,
|
||||
integration_session: AsyncSession,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Slow upstream header waits must not retain a checked-out DB connection."""
|
||||
key = ApiKey(
|
||||
hashed_key="proxy-pool-key",
|
||||
balance=1_000_000,
|
||||
refund_mint_url="http://primary:3338",
|
||||
refund_currency="sat",
|
||||
)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {"authorization": "Bearer test-key"}
|
||||
request.body = AsyncMock(return_value=json.dumps({"model": "test-model"}).encode())
|
||||
request.url.path = "/v1/chat/completions"
|
||||
request.state.request_id = "pool-hold-regression"
|
||||
|
||||
model = MagicMock()
|
||||
upstream = MagicMock()
|
||||
upstream.provider_type = "test"
|
||||
upstream.prepare_headers.return_value = {}
|
||||
|
||||
async def wait_for_headers(*args: object, **kwargs: object) -> Response:
|
||||
assert integration_engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return Response(status_code=200)
|
||||
|
||||
upstream.forward_request = AsyncMock(side_effect=wait_for_headers)
|
||||
|
||||
with (
|
||||
patch("routstr.proxy.get_candidates", return_value=[(model, upstream)]),
|
||||
patch("routstr.proxy.get_max_cost_for_model", AsyncMock(return_value=100)),
|
||||
patch(
|
||||
"routstr.proxy.calculate_discounted_max_cost",
|
||||
AsyncMock(return_value=100),
|
||||
),
|
||||
patch("routstr.proxy.check_token_balance"),
|
||||
patch("routstr.proxy.get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
):
|
||||
response = await proxy_module._proxy(
|
||||
request, "v1/chat/completions", integration_session
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
@@ -89,7 +89,12 @@ def _make_swap_mocks(
|
||||
def _wallet_router(primary_wallet: Mock, token_wallet: Mock) -> Callable[..., Mock]:
|
||||
"""Route get_wallet calls to the primary or foreign wallet mock by URL."""
|
||||
|
||||
def fake_get_wallet(mint_url: str, unit: str = "sat", load: bool = True) -> Mock:
|
||||
def fake_get_wallet(
|
||||
mint_url: str,
|
||||
unit: str = "sat",
|
||||
load: bool = True,
|
||||
**kwargs: object,
|
||||
) -> Mock:
|
||||
return primary_wallet if mint_url == PRIMARY_MINT else token_wallet
|
||||
|
||||
return fake_get_wallet
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
from contextlib import asynccontextmanager
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core.admin import get_transactions_api
|
||||
from routstr.core.db import CashuTransaction
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transactions_api_excludes_internal_sweep_claim_timestamp() -> None:
|
||||
transaction = CashuTransaction(
|
||||
token="cashu-token",
|
||||
amount=10,
|
||||
unit="sat",
|
||||
type="out",
|
||||
sweep_started_at=123,
|
||||
)
|
||||
count_result = MagicMock()
|
||||
count_result.one.return_value = 1
|
||||
transactions_result = MagicMock()
|
||||
transactions_result.all.return_value = [transaction]
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[count_result, transactions_result])
|
||||
|
||||
@asynccontextmanager
|
||||
async def create_session(): # type: ignore[no-untyped-def]
|
||||
yield session
|
||||
|
||||
with patch("routstr.core.admin.create_session", create_session):
|
||||
response = await get_transactions_api()
|
||||
|
||||
assert response["total"] == 1
|
||||
assert response["transactions"][0]["token"] == "cashu-token"
|
||||
assert "sweep_started_at" not in response["transactions"][0]
|
||||
@@ -66,6 +66,10 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected()
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(return_value=True),
|
||||
) as store,
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.token_mint_url",
|
||||
return_value="https://fallback-mint.test",
|
||||
),
|
||||
patch("routstr.upstream.auto_topup.create_session", return_value=session),
|
||||
):
|
||||
await _check_and_topup(_row())
|
||||
@@ -74,7 +78,7 @@ async def test_auto_topup_persists_before_sending_and_marks_success_collected()
|
||||
token="cashu-token",
|
||||
amount=50,
|
||||
unit="sat",
|
||||
mint_url="https://mint.test",
|
||||
mint_url="https://fallback-mint.test",
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="auto_topup",
|
||||
@@ -138,6 +142,12 @@ async def test_auto_topup_does_not_send_untracked_token() -> None:
|
||||
"routstr.upstream.auto_topup.store_cashu_transaction",
|
||||
AsyncMock(side_effect=RuntimeError("database unavailable")),
|
||||
),
|
||||
patch(
|
||||
"routstr.upstream.auto_topup.release_token_reservation",
|
||||
AsyncMock(),
|
||||
) as reclaim,
|
||||
):
|
||||
await _check_and_topup(_row())
|
||||
|
||||
reclaim.assert_awaited_once_with("cashu-token")
|
||||
provider.topup.assert_not_awaited()
|
||||
|
||||
@@ -534,6 +534,29 @@ async def test_topup_mint_unreachable_returns_503(error: Exception) -> None:
|
||||
assert exc_info.value.detail == "Cashu mint is unreachable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_unreachable_source_mint_explains_why_fallback_is_impossible() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.wallet import SourceMintConnectionError
|
||||
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
error = SourceMintConnectionError("Issuing Cashu mint is unreachable")
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.credit_balance", AsyncMock(side_effect=error)),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await topup_wallet_endpoint(
|
||||
cashu_token="cashuAtoken", key=key, session=session
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 503
|
||||
assert "cannot be redeemed at another mint" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_topup_already_spent_still_returns_400() -> None:
|
||||
"""Regression: the mint-unreachable short-circuit must not swallow the
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Real-DB coverage for db.balances_by_mint_and_unit.
|
||||
|
||||
Verifies the grouped liability query used by fetch_all_balances: it sums
|
||||
balances per (mint_url, unit), filters to the requested mints/units, excludes
|
||||
NULL mint/currency rows, and returns nothing for empty inputs.
|
||||
"""
|
||||
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlalchemy.pool import StaticPool
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import (
|
||||
ApiKey,
|
||||
balance_for_mint_and_unit,
|
||||
balances_by_mint_and_unit,
|
||||
)
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
return create_async_engine(
|
||||
"sqlite+aiosqlite://",
|
||||
poolclass=StaticPool,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def session() -> "AsyncGenerator[AsyncSession, None]":
|
||||
engine = _make_engine()
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
db_session = AsyncSession(engine, expire_on_commit=False)
|
||||
try:
|
||||
yield db_session
|
||||
finally:
|
||||
await db_session.close()
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _add_key(
|
||||
session: AsyncSession,
|
||||
hashed_key: str,
|
||||
balance: int,
|
||||
mint_url: str | None,
|
||||
currency: str | None,
|
||||
) -> None:
|
||||
session.add(
|
||||
ApiKey(
|
||||
hashed_key=hashed_key,
|
||||
balance=balance,
|
||||
refund_mint_url=mint_url,
|
||||
refund_currency=currency,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sums_and_groups_by_mint_and_unit(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 500, "http://m1", "sat")
|
||||
await _add_key(session, "c", 7000, "http://m1", "msat")
|
||||
await _add_key(session, "d", 200, "http://m2", "sat")
|
||||
|
||||
result = await balances_by_mint_and_unit(
|
||||
session, ["http://m1", "http://m2"], ["sat", "msat"]
|
||||
)
|
||||
|
||||
assert result[("http://m1", "sat")] == 1500
|
||||
assert result[("http://m1", "msat")] == 7000
|
||||
assert result[("http://m2", "sat")] == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filters_out_unrequested_mints_and_units(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://wanted", "sat")
|
||||
await _add_key(session, "b", 999, "http://other", "sat")
|
||||
await _add_key(session, "c", 888, "http://wanted", "usd")
|
||||
|
||||
result = await balances_by_mint_and_unit(session, ["http://wanted"], ["sat"])
|
||||
|
||||
assert result == {("http://wanted", "sat"): 1000}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_excludes_rows_with_null_mint_or_currency(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 4242, None, None)
|
||||
|
||||
result = await balances_by_mint_and_unit(session, ["http://m1"], ["sat"])
|
||||
|
||||
assert result == {("http://m1", "sat"): 1000}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scalar_balance_for_one_mint_and_unit(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 500, "http://m1", "sat")
|
||||
await _add_key(session, "c", 9000, "http://m1", "msat")
|
||||
await _add_key(session, "d", 700, "http://m2", "sat")
|
||||
|
||||
assert await balance_for_mint_and_unit(session, "http://m1", "sat") == 1500
|
||||
assert await balance_for_mint_and_unit(session, "http://missing", "sat") == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_inputs_return_empty_mapping(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
|
||||
assert await balances_by_mint_and_unit(session, [], ["sat"]) == {}
|
||||
assert await balances_by_mint_and_unit(session, ["http://m1"], []) == {}
|
||||
@@ -0,0 +1,85 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from routstr.core import db
|
||||
from routstr.core.db import create_db_engine
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_uses_validated_bounded_pool_settings(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: object
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_size", 12)
|
||||
monkeypatch.setattr(settings, "database_max_overflow", 3)
|
||||
monkeypatch.setattr(settings, "database_pool_timeout", 2.5)
|
||||
monkeypatch.setattr(settings, "database_pool_recycle", 900)
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
|
||||
engine = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/pool.db")
|
||||
try:
|
||||
assert engine.pool.size() == 12 # type: ignore[attr-defined]
|
||||
assert engine.pool._max_overflow == 3 # type: ignore[attr-defined]
|
||||
assert engine.pool._timeout == 2.5 # type: ignore[attr-defined]
|
||||
assert engine.pool._recycle == 900
|
||||
assert engine.pool._pre_ping is False
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_memory_sqlite_keeps_static_pool(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", True)
|
||||
engine = create_db_engine("sqlite+aiosqlite://")
|
||||
try:
|
||||
assert isinstance(engine.pool, StaticPool)
|
||||
assert engine.pool._pre_ping is True
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def test_non_sqlite_backend_enables_pre_ping_automatically(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
fake_engine = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(db, "create_async_engine", return_value=fake_engine) as factory,
|
||||
patch.object(db.event, "listen") as listen,
|
||||
):
|
||||
created = create_db_engine("postgresql+asyncpg://user:pass@db/node")
|
||||
|
||||
assert created is fake_engine
|
||||
assert factory.call_args.kwargs["pool_pre_ping"] is True
|
||||
assert listen.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_created_engine_warns_for_long_checkouts(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: object
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_hold_warn_seconds", 0.0)
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
first = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/first.db")
|
||||
second = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/second.db")
|
||||
|
||||
try:
|
||||
with patch.object(db.logger, "warning") as warning:
|
||||
async with first.connect() as connection:
|
||||
await connection.exec_driver_sql("SELECT 1")
|
||||
async with second.connect() as connection:
|
||||
await connection.exec_driver_sql("SELECT 1")
|
||||
|
||||
assert warning.call_count == 2
|
||||
assert all(
|
||||
call.kwargs["extra"]["threshold_seconds"] == 0.0
|
||||
for call in warning.call_args_list
|
||||
)
|
||||
finally:
|
||||
await first.dispose()
|
||||
await second.dispose()
|
||||
@@ -1,5 +1,5 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
@@ -13,9 +13,19 @@ from routstr import wallet
|
||||
from routstr.core import db
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _session_context(session: Mock) -> AsyncIterator[Mock]:
|
||||
yield session
|
||||
class _SessionContext:
|
||||
def __init__(self, session: Mock) -> None:
|
||||
self.session = session
|
||||
|
||||
async def __aenter__(self) -> Mock:
|
||||
return self.session
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _session_context(session: Mock) -> _SessionContext:
|
||||
return _SessionContext(session)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -47,7 +57,7 @@ async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
@@ -57,6 +67,10 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
payout_wallet = Mock()
|
||||
events: list[str] = []
|
||||
|
||||
async def prepare(*_args: object) -> Mock:
|
||||
events.append("prepare")
|
||||
return payout_wallet
|
||||
|
||||
async def checkpoint(*_args: object) -> bool:
|
||||
events.append("checkpoint")
|
||||
return True
|
||||
@@ -77,18 +91,92 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(side_effect=prepare)),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
assert events == ["checkpoint", "send", "complete"]
|
||||
assert events == ["prepare", "checkpoint", "send", "complete"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_preparation_failure_does_not_checkpoint() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
checkpoint = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", checkpoint),
|
||||
patch(
|
||||
"routstr.wallet.get_wallet",
|
||||
AsyncMock(side_effect=RuntimeError("wallet unavailable")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
checkpoint.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
send = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch(
|
||||
"routstr.wallet.db.reset_routstr_fee",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", send),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
send.assert_not_awaited()
|
||||
warning.assert_called_once_with("Routstr fee payout was already claimed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -107,7 +195,9 @@ async def test_fee_payout_does_not_retry_an_unresolved_checkpoint() -> None:
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock()) as checkpoint,
|
||||
patch("routstr.wallet.get_wallet", AsyncMock()) as get_wallet,
|
||||
@@ -141,7 +231,9 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
|
||||
@@ -158,3 +250,139 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non
|
||||
|
||||
complete.assert_not_awaited()
|
||||
critical.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_cancellation_during_send_alerts_and_propagates() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
complete = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch("routstr.wallet.asyncio.sleep", AsyncMock(return_value=None)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch(
|
||||
"routstr.wallet.raw_send_to_lnurl",
|
||||
AsyncMock(side_effect=asyncio.CancelledError()),
|
||||
),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
complete.assert_not_awaited()
|
||||
critical.assert_called_once()
|
||||
assert critical.call_args.args[0] == (
|
||||
"Routstr fee payout outcome is unknown; manual reconciliation required"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure_site", ["session", "completion"])
|
||||
async def test_fee_payout_completion_failures_use_sent_checkpoint_alert(
|
||||
failure_site: str,
|
||||
) -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
completion = AsyncMock()
|
||||
if failure_site == "session":
|
||||
create_session = Mock(
|
||||
side_effect=[
|
||||
_session_context(session),
|
||||
_session_context(session),
|
||||
RuntimeError("pool unavailable"),
|
||||
]
|
||||
)
|
||||
else:
|
||||
create_session = Mock(return_value=_session_context(session))
|
||||
completion.side_effect = RuntimeError("checkpoint unavailable")
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", create_session),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", completion),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(return_value=5)),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
critical.assert_called_once()
|
||||
assert critical.call_args.args[0] == (
|
||||
"Routstr fee payout sent but checkpoint was not completed"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_releases_db_connection_during_send(tmp_path: object) -> None:
|
||||
"""With pool_size=1, the payout must not hold a connection while the
|
||||
external LNURL send is in flight, or the completion step would starve."""
|
||||
engine = create_async_engine(
|
||||
f"sqlite+aiosqlite:///{tmp_path}/payout.db", pool_size=1, max_overflow=0
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
async with AsyncSession(engine) as session:
|
||||
session.add(db.RoutstrFee(id=1, accumulated_msats=5_000_000))
|
||||
await session.commit()
|
||||
|
||||
@asynccontextmanager
|
||||
async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
|
||||
async def send(*_args: object, **_kwargs: object) -> int:
|
||||
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return 5
|
||||
|
||||
try:
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", create_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
fee = await db.get_routstr_fee(session)
|
||||
assert fee.payout_in_progress_msats == 0
|
||||
assert fee.total_paid_msats == 5_000_000
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
@@ -37,7 +37,7 @@ def test_fresh_node_migrates_fee_payout_schema_to_head(tmp_path: Path) -> None:
|
||||
"payout_in_progress_msats, payout_started_at FROM routstr_fees"
|
||||
).fetchone()
|
||||
|
||||
assert version == ("9c4d8e2f1a6b",)
|
||||
assert version == ("bf76270b66c4",)
|
||||
assert {
|
||||
"id",
|
||||
"accumulated_msats",
|
||||
|
||||
@@ -1,11 +1,33 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncGenerator, Generator
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.wallet import fetch_all_balances
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_balance_fetch_state() -> Generator[None, None, None]:
|
||||
from routstr import wallet
|
||||
|
||||
wallet._balance_fetch_failures.clear()
|
||||
wallet._balance_fetch_locks.clear()
|
||||
wallet._mint_supported_units.clear()
|
||||
wallet._MintRateGuard._guards.clear()
|
||||
yield
|
||||
wallet._balance_fetch_failures.clear()
|
||||
wallet._balance_fetch_locks.clear()
|
||||
wallet._mint_supported_units.clear()
|
||||
wallet._MintRateGuard._guards.clear()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _fake_session(): # type: ignore[no-untyped-def]
|
||||
yield MagicMock()
|
||||
@@ -23,11 +45,13 @@ def _patches( # type: ignore[no-untyped-def]
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
AsyncMock(side_effect=lambda proofs, wallet, **kwargs: proofs),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit",
|
||||
AsyncMock(return_value=user_balance_msats),
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(
|
||||
return_value={("http://primary:3338", "sat"): user_balance_msats}
|
||||
),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
]
|
||||
@@ -38,8 +62,9 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None:
|
||||
"""With empty cashu_mints, balances are still fetched for primary_mint."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with patch.object(settings, "cashu_mints", []), patch.object(
|
||||
settings, "primary_mint", "http://primary:3338"
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
):
|
||||
for p in _patches(proof_amount=1000):
|
||||
p.start()
|
||||
@@ -54,13 +79,314 @@ async def test_fetch_all_balances_falls_back_to_primary_mint() -> None:
|
||||
assert total_wallet == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_uses_units_advertised_by_mint() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch(
|
||||
"routstr.wallet._get_supported_mint_units",
|
||||
AsyncMock(return_value=["sat"]),
|
||||
) as supported_units,
|
||||
):
|
||||
for p in _patches(proof_amount=1000):
|
||||
p.start()
|
||||
try:
|
||||
details, *_ = await fetch_all_balances()
|
||||
finally:
|
||||
patch.stopall()
|
||||
|
||||
supported_units.assert_awaited_once_with("http://mint:3338")
|
||||
assert [detail["unit"] for detail in details] == ["sat"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unit_discovery_failure_returns_structured_balance_error() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
get_wallet = AsyncMock()
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch(
|
||||
"routstr.wallet._get_supported_mint_units",
|
||||
AsyncMock(side_effect=httpx.ConnectError("mint unavailable")),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
):
|
||||
details, *_ = await fetch_all_balances()
|
||||
|
||||
assert details[0]["unit"] == settings.primary_mint_unit
|
||||
assert details[0]["error_code"] == "unreachable"
|
||||
assert details[0]["retry_after_seconds"] > 0
|
||||
get_wallet.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_supported_mint_units_come_from_active_keysets() -> None:
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import _get_supported_mint_units
|
||||
|
||||
# Cashu versions/mints may deserialize keyset units as either strings or
|
||||
# Unit enum-like objects. Both representations must be accepted.
|
||||
sat = MagicMock(active=True, unit="sat")
|
||||
msat = MagicMock(active=False, unit="msat")
|
||||
usd = MagicMock(active=True)
|
||||
usd.unit.name = "usd"
|
||||
wallet = MagicMock()
|
||||
wallet._get_keysets = AsyncMock(return_value=[usd, msat, sat])
|
||||
|
||||
with (
|
||||
patch.object(settings, "primary_mint_unit", "sat"),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=wallet)),
|
||||
):
|
||||
units = await _get_supported_mint_units("http://mint:3338")
|
||||
cached_units = await _get_supported_mint_units("http://mint:3338")
|
||||
|
||||
assert units == ["sat", "usd"]
|
||||
assert cached_units == units
|
||||
wallet._get_keysets.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_backs_off_after_connection_failure() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
get_wallet = AsyncMock(side_effect=httpx.ConnectError("mint unavailable"))
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.wallet.time.monotonic", return_value=10),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
first = await fetch_all_balances(units=["sat"])
|
||||
second = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert first[0][0]["error"] == "mint unavailable"
|
||||
assert first[0][0]["error_code"] == "unreachable"
|
||||
assert first[0][0]["retry_after_seconds"] == 60
|
||||
assert second[0][0]["error"] == "mint unavailable"
|
||||
assert second[0][0]["error_code"] == "unreachable"
|
||||
assert get_wallet.await_count == 1
|
||||
warning.assert_called_once()
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.wallet.time.monotonic", return_value=71),
|
||||
patch("routstr.wallet.logger.warning"),
|
||||
):
|
||||
await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert get_wallet.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_reports_rate_limit_status() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
request = httpx.Request("GET", "http://mint:3338/v1/keysets")
|
||||
response = httpx.Response(429, request=request, headers={"Retry-After": "45"})
|
||||
error = httpx.HTTPStatusError("rate limited", request=request, response=response)
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(side_effect=error)),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert details[0]["error_code"] == "rate_limited"
|
||||
assert details[0]["retry_after_seconds"] == 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_failure_applies_mint_cooldown_to_other_units() -> None:
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import _mint_cooldown_remaining
|
||||
|
||||
mint = "http://mint:3338"
|
||||
get_wallet = AsyncMock(side_effect=httpx.ConnectError("mint unavailable"))
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", [mint]),
|
||||
patch.object(settings, "primary_mint", mint),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.wallet.time.monotonic", return_value=10),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat", "msat"])
|
||||
cooldown = _mint_cooldown_remaining(mint)
|
||||
|
||||
assert get_wallet.await_count == 1
|
||||
assert warning.call_count == 1
|
||||
assert cooldown == 60
|
||||
assert details[0]["error"] == "mint unavailable"
|
||||
assert details[0]["error_code"] == "unreachable"
|
||||
assert details[1]["error"] == "Mint is unreachable"
|
||||
assert details[1]["error_code"] == "unreachable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_closes_db_session_before_concurrent_mint_io() -> None:
|
||||
"""Slow mint checks must never run while the balance DB session is open."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
session_open = False
|
||||
mint_calls = 0
|
||||
|
||||
@asynccontextmanager
|
||||
async def tracked_session(): # type: ignore[no-untyped-def]
|
||||
nonlocal session_open
|
||||
session_open = True
|
||||
try:
|
||||
yield MagicMock()
|
||||
finally:
|
||||
session_open = False
|
||||
|
||||
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
|
||||
nonlocal mint_calls
|
||||
assert session_open is False
|
||||
mint_calls += 1
|
||||
await asyncio.sleep(0)
|
||||
return proofs
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://one:3338", "http://two:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://one:3338"),
|
||||
patch("routstr.wallet.db.create_session", tracked_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(return_value={}),
|
||||
create=True,
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=1)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=slow_filter),
|
||||
),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat", "msat"])
|
||||
|
||||
assert mint_calls == 4
|
||||
assert all("error" not in detail for detail in details)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_bounds_parallel_mint_checks() -> None:
|
||||
"""A slow mint fleet cannot create an unbounded external-I/O fan-out."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
active = 0
|
||||
peak = 0
|
||||
|
||||
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
|
||||
nonlocal active, peak
|
||||
active += 1
|
||||
peak = max(peak, active)
|
||||
await asyncio.sleep(0.01)
|
||||
active -= 1
|
||||
return proofs
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
settings,
|
||||
"cashu_mints",
|
||||
[f"http://mint-{index}:3338" for index in range(8)],
|
||||
),
|
||||
patch.object(settings, "primary_mint", ""),
|
||||
patch.object(settings, "mint_operation_concurrency", 2),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(return_value={}),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=slow_filter),
|
||||
),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert len(details) == 8
|
||||
assert peak == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slow_mints_do_not_exhaust_a_single_connection_pool(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Concurrent slow balance refreshes release the sole DB connection promptly."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
engine = create_async_engine(
|
||||
f"sqlite+aiosqlite:///{tmp_path / 'pool-pressure.db'}",
|
||||
pool_size=1,
|
||||
max_overflow=0,
|
||||
pool_timeout=0.2,
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
@asynccontextmanager
|
||||
async def single_pool_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
|
||||
async def slow_filter(proofs, wallet): # type: ignore[no-untyped-def]
|
||||
await asyncio.sleep(0.3)
|
||||
return proofs
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://slow:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://slow:3338"),
|
||||
patch.object(settings, "mint_operation_concurrency", 1),
|
||||
patch("routstr.wallet.db.create_session", single_pool_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=slow_filter),
|
||||
),
|
||||
):
|
||||
results = await asyncio.gather(
|
||||
*(fetch_all_balances(units=["sat"]) for _ in range(6))
|
||||
)
|
||||
|
||||
assert all("error" not in result[0][0] for result in results)
|
||||
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_reports_liability_when_wallet_is_empty() -> None:
|
||||
"""An empty wallet must not hide outstanding user liabilities."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with patch.object(settings, "cashu_mints", []), patch.object(
|
||||
settings, "primary_mint", "http://primary:3338"
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
):
|
||||
for p in _patches(proof_amount=0, user_balance_msats=5000):
|
||||
p.start()
|
||||
@@ -84,9 +410,10 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None:
|
||||
"""primary_mint already in cashu_mints is not inspected twice."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with patch.object(
|
||||
settings, "cashu_mints", ["http://primary:3338"]
|
||||
), patch.object(settings, "primary_mint", "http://primary:3338"):
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://primary:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
):
|
||||
for p in _patches(proof_amount=1000):
|
||||
p.start()
|
||||
try:
|
||||
@@ -98,3 +425,58 @@ async def test_fetch_all_balances_no_duplicate_primary_mint() -> None:
|
||||
|
||||
assert [d["mint_url"] for d in details] == ["http://primary:3338"]
|
||||
assert total_wallet == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_all_balances_degrades_when_liability_read_fails() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(side_effect=RuntimeError("db pool exhausted")),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=1000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
):
|
||||
details, total_wallet, total_user, owner = await fetch_all_balances(
|
||||
units=["sat"]
|
||||
)
|
||||
|
||||
assert details[0]["error"] == "db pool exhausted"
|
||||
assert details[0]["wallet_balance"] == 1000
|
||||
assert details[0]["user_balance"] == 0
|
||||
assert details[0]["owner_balance"] == 0
|
||||
assert (total_wallet, total_user, owner) == (1000, 0, 0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_liability_error_keeps_more_specific_mint_error() -> None:
|
||||
from routstr.core.settings import settings
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch(
|
||||
"routstr.wallet.db.balances_by_mint_and_unit",
|
||||
AsyncMock(side_effect=RuntimeError("db pool exhausted")),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.get_wallet",
|
||||
AsyncMock(side_effect=RuntimeError("mint down")),
|
||||
),
|
||||
):
|
||||
details, *_ = await fetch_all_balances(units=["sat"])
|
||||
|
||||
assert details[0]["error"] == "mint down"
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from cashu.core.base import Proof
|
||||
|
||||
from routstr.lightning import (
|
||||
_invoice_settlement_locks,
|
||||
_is_outputs_already_signed,
|
||||
_mint_invoice_quote,
|
||||
check_invoice_payment,
|
||||
)
|
||||
from routstr.wallet import Wallet
|
||||
|
||||
|
||||
def _invoice(**overrides: object) -> SimpleNamespace:
|
||||
values = {
|
||||
"id": "invoice-1",
|
||||
"payment_hash": "quote-1",
|
||||
"amount_sats": 100,
|
||||
"purpose": "create",
|
||||
"status": "pending",
|
||||
"paid_at": None,
|
||||
"api_key_hash": None,
|
||||
"mint_url": "http://mint:3338",
|
||||
"balance_limit": None,
|
||||
"balance_limit_reset": None,
|
||||
"validity_date": None,
|
||||
}
|
||||
values.update(overrides)
|
||||
return SimpleNamespace(**values)
|
||||
|
||||
|
||||
def _proof(amount: int, mint_id: str, *, reserved: bool = False) -> Proof:
|
||||
return Proof(amount=amount, mint_id=mint_id, reserved=reserved)
|
||||
|
||||
|
||||
def _recovery_wallet(
|
||||
error: Exception,
|
||||
*,
|
||||
proofs_before: list[Proof] | None = None,
|
||||
proofs_after: list[Proof] | None = None,
|
||||
) -> Mock:
|
||||
async def load_proofs(*, reload: bool) -> None:
|
||||
if wallet.load_proofs.await_count >= 2 and proofs_after is not None:
|
||||
wallet.proofs = list(proofs_after)
|
||||
|
||||
wallet = Mock(
|
||||
mint=AsyncMock(side_effect=error),
|
||||
keysets={"keyset-1": Mock()},
|
||||
restore_tokens_for_keyset=AsyncMock(),
|
||||
load_proofs=AsyncMock(side_effect=load_proofs),
|
||||
proofs=list(proofs_before or []),
|
||||
)
|
||||
return wallet
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_recovers_quote_linked_outputs_already_signed() -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("Mint Error: outputs have already been signed before (Code: 11003)"),
|
||||
proofs_after=[_proof(100, "quote-1")],
|
||||
)
|
||||
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
wallet.restore_tokens_for_keyset.assert_awaited_once_with(
|
||||
"keyset-1", to=1, batch=25
|
||||
)
|
||||
assert wallet.load_proofs.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_accepts_preloaded_quote_linked_proofs() -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("must not mint"),
|
||||
proofs_before=[_proof(64, "quote-1"), _proof(36, "quote-1")],
|
||||
)
|
||||
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
wallet.mint.assert_not_awaited()
|
||||
wallet.restore_tokens_for_keyset.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_does_not_accept_unrelated_11003_text() -> None:
|
||||
invoice = _invoice()
|
||||
error = Exception("backend request 11003 failed")
|
||||
wallet = _recovery_wallet(error)
|
||||
|
||||
with pytest.raises(Exception) as caught:
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
assert caught.value is error
|
||||
wallet.restore_tokens_for_keyset.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_installed_cashu_error_shape_recognizes_realistic_11003_phrase() -> None:
|
||||
request = httpx.Request("POST", "http://mint:3338/v1/mint/bolt11")
|
||||
response = httpx.Response(
|
||||
400,
|
||||
request=request,
|
||||
json={"detail": "outputs have already been signed before", "code": 11003},
|
||||
)
|
||||
|
||||
with pytest.raises(Exception) as caught:
|
||||
Wallet.raise_on_error_request(response)
|
||||
|
||||
assert _is_outputs_already_signed(caught.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("recovered", [0, 99])
|
||||
async def test_invoice_mint_rejects_empty_or_short_quote_recovery(
|
||||
recovered: int,
|
||||
) -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("Mint Error: outputs already signed (Code: 11003)"),
|
||||
proofs_after=[_proof(recovered, "quote-1")] if recovered else [],
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="expected at least 100"):
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_mint_rejects_unrelated_concurrent_balance_growth() -> None:
|
||||
invoice = _invoice()
|
||||
wallet = _recovery_wallet(
|
||||
Exception("Mint Error: outputs already signed (Code: 11003)"),
|
||||
proofs_after=[_proof(10_000, "different-quote")],
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="quote-linked recovery returned 0"):
|
||||
await _mint_invoice_quote(wallet, invoice) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_pending_invoice_is_not_minted() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice(status="expired")
|
||||
session = AsyncMock()
|
||||
|
||||
with patch("routstr.lightning.get_wallet", AsyncMock()) as get_wallet:
|
||||
await check_invoice_payment(invoice, session) # type: ignore[arg-type]
|
||||
|
||||
get_wallet.assert_not_awaited()
|
||||
session.commit.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ambiguous_invoice_mint_timeout_does_not_expose_paid() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice()
|
||||
session = AsyncMock()
|
||||
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True)))
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch(
|
||||
"routstr.lightning._mint_invoice_quote",
|
||||
AsyncMock(side_effect=httpx.TimeoutException("response lost")),
|
||||
),
|
||||
patch("routstr.lightning._reload_invoice_view", AsyncMock()),
|
||||
):
|
||||
await check_invoice_payment(invoice, session) # type: ignore[arg-type]
|
||||
|
||||
assert invoice.status == "pending"
|
||||
session.rollback.assert_not_awaited()
|
||||
# One commit closes the initial read transaction before external I/O.
|
||||
session.commit.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_invoice_checks_finalize_once_in_process() -> None:
|
||||
_invoice_settlement_locks.clear()
|
||||
invoice = _invoice()
|
||||
session = AsyncMock()
|
||||
wallet = Mock(get_mint_quote=AsyncMock(return_value=Mock(paid=True)))
|
||||
|
||||
async def refresh(obj: SimpleNamespace) -> None:
|
||||
return None
|
||||
|
||||
session.refresh = AsyncMock(side_effect=refresh)
|
||||
|
||||
@asynccontextmanager
|
||||
async def owned_session() -> AsyncIterator[AsyncMock]:
|
||||
yield AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.lightning.get_wallet", AsyncMock(return_value=wallet)),
|
||||
patch("routstr.lightning.create_session", owned_session),
|
||||
patch("routstr.lightning._mint_invoice_quote", AsyncMock()),
|
||||
patch(
|
||||
"routstr.lightning._finalize_invoice_settlement",
|
||||
AsyncMock(return_value=(True, "b" * 64)),
|
||||
) as finalize,
|
||||
):
|
||||
await asyncio.gather(
|
||||
check_invoice_payment(invoice, session), # type: ignore[arg-type]
|
||||
check_invoice_payment(invoice, session), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
assert invoice.status == "paid"
|
||||
finalize.assert_awaited_once()
|
||||
@@ -0,0 +1,47 @@
|
||||
import os
|
||||
import sqlite3
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _run_alembic(root: Path, database_url: str, command: str, revision: str) -> None:
|
||||
env = os.environ.copy()
|
||||
env["DATABASE_URL"] = database_url
|
||||
subprocess.run(
|
||||
[sys.executable, "-m", "alembic", command, revision],
|
||||
cwd=root,
|
||||
env=env,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
|
||||
def _lightning_invoice_columns(database_path: Path) -> set[str]:
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
return {
|
||||
row[1]
|
||||
for row in connection.execute("PRAGMA table_info(lightning_invoices)")
|
||||
}
|
||||
|
||||
|
||||
def test_mint_url_migration_upgrades_and_downgrades_from_main_head(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
root = Path(__file__).resolve().parents[2]
|
||||
database_path = tmp_path / "mint-url-migration.db"
|
||||
database_url = f"sqlite+aiosqlite:///{database_path}"
|
||||
previous_head = "aa50fde387a2"
|
||||
|
||||
_run_alembic(root, database_url, "upgrade", previous_head)
|
||||
assert "mint_url" not in _lightning_invoice_columns(database_path)
|
||||
|
||||
_run_alembic(root, database_url, "upgrade", "bf76270b66c4")
|
||||
assert "mint_url" in _lightning_invoice_columns(database_path)
|
||||
|
||||
_run_alembic(root, database_url, "downgrade", previous_head)
|
||||
assert "mint_url" not in _lightning_invoice_columns(database_path)
|
||||
|
||||
_run_alembic(root, database_url, "upgrade", "head")
|
||||
assert "mint_url" in _lightning_invoice_columns(database_path)
|
||||
@@ -7,7 +7,21 @@ os.environ["UPSTREAM_BASE_URL"] = "http://test"
|
||||
os.environ["UPSTREAM_API_KEY"] = "test"
|
||||
|
||||
from routstr.core.settings import settings # noqa: E402
|
||||
from routstr.payment.helpers import get_max_cost_for_model # noqa: E402
|
||||
from routstr.payment.helpers import ( # noqa: E402
|
||||
apply_mint_fee_allowance,
|
||||
get_max_cost_for_model,
|
||||
)
|
||||
|
||||
|
||||
def test_mint_fee_allowance_reserves_five_percent_fallback_headroom() -> None:
|
||||
# Interim policy: Routstr may pay hidden cross-mint Lightning fees when a
|
||||
# trusted-mint fallback is required.
|
||||
assert apply_mint_fee_allowance(124_886) == 118_642
|
||||
|
||||
|
||||
def test_mint_fee_allowance_never_drops_below_minimum() -> None:
|
||||
with patch.object(settings, "min_request_msat", 100):
|
||||
assert apply_mint_fee_allowance(50) == 100
|
||||
|
||||
|
||||
async def test_get_max_cost_for_model_known() -> None:
|
||||
|
||||
@@ -59,23 +59,29 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
|
||||
get_wallet = AsyncMock(return_value=MagicMock())
|
||||
raw_send = AsyncMock(return_value=1000)
|
||||
|
||||
with patch.object(settings, "cashu_mints", []), patch.object(
|
||||
settings, "primary_mint", "http://primary:3338"
|
||||
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
|
||||
settings, "payout_interval_seconds", _INTERVAL
|
||||
), patch.object(settings, "min_payout_sat", 10), patch(
|
||||
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
|
||||
), patch("routstr.wallet.db.create_session", _fake_session), patch(
|
||||
"routstr.wallet.get_wallet", get_wallet
|
||||
), patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
), patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
), patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0)
|
||||
), patch("routstr.wallet.raw_send_to_lnurl", raw_send):
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.tld"),
|
||||
patch.object(settings, "payout_interval_seconds", _INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.total_user_liability",
|
||||
AsyncMock(return_value=0),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
@@ -84,6 +90,58 @@ async def test_periodic_payout_includes_primary_mint_not_in_cashu_mints() -> Non
|
||||
assert raw_send.await_count >= 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_releases_session_before_slow_mint_send() -> None:
|
||||
"""The DB connection is returned before the external LNURL call starts."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
session_open = False
|
||||
sends_completed = 0
|
||||
|
||||
@asynccontextmanager
|
||||
async def tracked_session(): # type: ignore[no-untyped-def]
|
||||
nonlocal session_open
|
||||
session_open = True
|
||||
try:
|
||||
yield MagicMock()
|
||||
finally:
|
||||
session_open = False
|
||||
|
||||
async def raw_send(*args: object, **kwargs: object) -> int:
|
||||
nonlocal sends_completed
|
||||
assert session_open is False
|
||||
sends_completed += 1
|
||||
return 1000
|
||||
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", []),
|
||||
patch.object(settings, "primary_mint", "http://primary:3338"),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.tld"),
|
||||
patch.object(settings, "payout_interval_seconds", _INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
|
||||
patch("routstr.wallet.db.create_session", tracked_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.total_user_liability",
|
||||
AsyncMock(return_value=0),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(side_effect=raw_send)),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
assert sends_completed == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_isolates_failing_mint() -> None:
|
||||
"""A failing mint does not prevent payout for the other mints."""
|
||||
@@ -97,23 +155,29 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
|
||||
get_wallet = AsyncMock(side_effect=_get_wallet)
|
||||
raw_send = AsyncMock(return_value=1000)
|
||||
|
||||
with patch.object(
|
||||
settings, "cashu_mints", ["http://bad:3338", "http://good:3338"]
|
||||
), patch.object(settings, "primary_mint", "http://good:3338"), patch.object(
|
||||
settings, "receive_ln_address", "owner@ln.tld"
|
||||
), patch.object(settings, "payout_interval_seconds", _INTERVAL), patch.object(
|
||||
settings, "min_payout_sat", 10
|
||||
), patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()), patch(
|
||||
"routstr.wallet.db.create_session", _fake_session
|
||||
), patch("routstr.wallet.get_wallet", get_wallet), patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
), patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
), patch(
|
||||
"routstr.wallet.db.balances_for_mint_and_unit", AsyncMock(return_value=0)
|
||||
), patch("routstr.wallet.raw_send_to_lnurl", raw_send):
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://bad:3338", "http://good:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://good:3338"),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.tld"),
|
||||
patch.object(settings, "payout_interval_seconds", _INTERVAL),
|
||||
patch.object(settings, "min_payout_sat", 10),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
|
||||
patch("routstr.wallet.db.create_session", _fake_session),
|
||||
patch("routstr.wallet.get_wallet", get_wallet),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.total_user_liability",
|
||||
AsyncMock(return_value=0),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
@@ -128,27 +192,39 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_handles_session_creation_failure() -> None:
|
||||
"""A db.create_session failure is logged and the payout loop continues."""
|
||||
"""A db.create_session failure is logged per mint/unit and the loop continues."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
create_session = MagicMock(side_effect=RuntimeError("db unavailable"))
|
||||
logger = MagicMock()
|
||||
|
||||
with patch.object(settings, "cashu_mints", ["http://mint:3338"]), patch.object(
|
||||
settings, "primary_mint", "http://mint:3338"
|
||||
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
|
||||
settings, "payout_interval_seconds", _INTERVAL
|
||||
), patch(
|
||||
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
|
||||
), patch(
|
||||
"routstr.wallet.db.create_session", create_session
|
||||
), patch("routstr.wallet.logger", logger):
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.tld"),
|
||||
patch.object(settings, "payout_interval_seconds", _INTERVAL),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
|
||||
patch("routstr.wallet.db.create_session", create_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
patch("routstr.wallet.logger", logger),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
create_session.assert_called_once()
|
||||
logger.error.assert_called_once()
|
||||
# The liability session is opened per mint/unit (sat + msat), and each
|
||||
# DB failure retains the cycle-specific alert wording while remaining
|
||||
# isolated to its own iteration.
|
||||
assert create_session.call_count == 2
|
||||
assert logger.error.call_count == 2
|
||||
message = logger.error.call_args.args[0]
|
||||
extra = logger.error.call_args.kwargs["extra"]
|
||||
assert message == "Error in periodic payout cycle: RuntimeError"
|
||||
assert extra == {"error": "db unavailable"}
|
||||
assert extra["error"] == "db unavailable"
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
from collections.abc import AsyncIterator
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_closes_request_session_before_returning_response() -> None:
|
||||
"""Route completion must release DB resources before response delivery."""
|
||||
request = MagicMock()
|
||||
request.method = "GET"
|
||||
request.headers = {"accept": "application/json"}
|
||||
request.url.path = "/not-an-api-route"
|
||||
request.state.request_id = "test-request"
|
||||
session = AsyncMock()
|
||||
|
||||
response = await proxy_module.proxy(request, "not-an-api-route", session=session)
|
||||
|
||||
assert response.status_code == 404
|
||||
session.close.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_session_is_closed_before_first_stream_chunk() -> None:
|
||||
request = MagicMock()
|
||||
session = AsyncMock()
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
session.close.assert_awaited_once()
|
||||
yield b"chunk"
|
||||
|
||||
upstream_response = StreamingResponse(stream())
|
||||
with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)):
|
||||
response = await proxy_module.proxy(
|
||||
request, "v1/chat/completions", session=session
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
chunks = [chunk async for chunk in response.body_iterator]
|
||||
assert chunks == [b"chunk"]
|
||||
@@ -1,4 +1,6 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
@@ -7,6 +9,7 @@ from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr import wallet
|
||||
from routstr.core.db import CashuTransaction
|
||||
from routstr.wallet import refund_sweep_once
|
||||
|
||||
@@ -39,6 +42,43 @@ async def _load(
|
||||
return {row.token: row for row in result.all()}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_releases_db_session_during_token_redemption(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="eligible", amount=1, unit="sat", type="out", created_at=800
|
||||
),
|
||||
)
|
||||
session_open = False
|
||||
|
||||
@asynccontextmanager
|
||||
async def tracked_session() -> AsyncIterator[AsyncSession]:
|
||||
nonlocal session_open
|
||||
async with session_factory() as session:
|
||||
session_open = True
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session_open = False
|
||||
|
||||
async def receive_token(token: str) -> None:
|
||||
assert token == "eligible"
|
||||
assert session_open is False
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", tracked_session),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=receive_token)),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
assert (await _load(session_factory))["eligible"].swept is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
@@ -90,14 +130,17 @@ async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens(
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("error", "collected"),
|
||||
("error", "collected", "claim_started_at"),
|
||||
[
|
||||
(RuntimeError("token already spent"), True),
|
||||
(RuntimeError("mint unavailable"), False),
|
||||
(RuntimeError("token already spent"), True, None),
|
||||
(RuntimeError("mint unavailable"), False, 1000),
|
||||
],
|
||||
)
|
||||
async def test_refund_sweep_records_terminal_but_not_transient_failures(
|
||||
session_factory: async_sessionmaker[AsyncSession], error: Exception, collected: bool
|
||||
async def test_refund_sweep_records_spent_and_unknown_outcomes_safely(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
error: Exception,
|
||||
collected: bool,
|
||||
claim_started_at: int | None,
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
@@ -116,3 +159,236 @@ async def test_refund_sweep_records_terminal_but_not_transient_failures(
|
||||
refund = (await _load(session_factory))["refund"]
|
||||
assert refund.collected is collected
|
||||
assert refund.swept is False
|
||||
assert refund.sweep_started_at == claim_started_at
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_spend_failure_retains_claim_and_stale_retry_records_sweep(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="post-spend-failure",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(
|
||||
side_effect=wallet.TokenConsumedError(
|
||||
"Mint on primary failed after successful melt"
|
||||
)
|
||||
),
|
||||
),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
retained = (await _load(session_factory))["post-spend-failure"]
|
||||
assert retained.swept is False
|
||||
assert retained.collected is False
|
||||
assert retained.sweep_started_at == 1000
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1300),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(side_effect=RuntimeError("token already spent")),
|
||||
),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
recovered = (await _load(session_factory))["post-spend-failure"]
|
||||
assert recovered.swept is True
|
||||
assert recovered.collected is False
|
||||
assert recovered.sweep_started_at is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_retains_claim_on_cancellation_during_redemption(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="cancelled", amount=1, unit="sat", type="out", created_at=800
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(side_effect=asyncio.CancelledError()),
|
||||
),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await refund_sweep_once()
|
||||
|
||||
refund = (await _load(session_factory))["cancelled"]
|
||||
assert refund.swept is False
|
||||
assert refund.sweep_started_at == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_checkpoint_failure_retains_claim_and_stale_retry_records_sweep(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="checkpoint-failure",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
),
|
||||
)
|
||||
real_set_state = wallet._set_refund_sweep_state
|
||||
|
||||
async def fail_swept_checkpoint(
|
||||
refund_id: str,
|
||||
*,
|
||||
predicates: tuple[object, ...] = (),
|
||||
**values: object,
|
||||
) -> int:
|
||||
if values.get("swept") is True:
|
||||
raise RuntimeError("checkpoint unavailable")
|
||||
return await real_set_state(refund_id, predicates=predicates, **values)
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token", AsyncMock(return_value=(1, "sat", "mint"))
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet._set_refund_sweep_state",
|
||||
side_effect=fail_swept_checkpoint,
|
||||
),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
retained = (await _load(session_factory))["checkpoint-failure"]
|
||||
assert retained.swept is False
|
||||
assert retained.collected is False
|
||||
assert retained.sweep_started_at == 1000
|
||||
critical.assert_called_once()
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1300),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(side_effect=RuntimeError("token already spent")),
|
||||
),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
recovered = (await _load(session_factory))["checkpoint-failure"]
|
||||
assert recovered.swept is True
|
||||
assert recovered.collected is False
|
||||
assert recovered.sweep_started_at is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("redemption_succeeds", [True, False])
|
||||
async def test_expired_worker_cannot_overwrite_or_release_newer_claim(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
redemption_succeeds: bool,
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="reclaimed",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
),
|
||||
)
|
||||
|
||||
async def replace_claim(_token: str) -> tuple[int, str, str]:
|
||||
async with session_factory() as session:
|
||||
result = await session.exec(
|
||||
select(CashuTransaction).where(CashuTransaction.token == "reclaimed")
|
||||
)
|
||||
transaction = result.one()
|
||||
transaction.sweep_started_at = 1100
|
||||
session.add(transaction)
|
||||
await session.commit()
|
||||
if not redemption_succeeds:
|
||||
raise RuntimeError("mint unavailable")
|
||||
return (1, "sat", "mint")
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=replace_claim)),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
reclaimed = (await _load(session_factory))["reclaimed"]
|
||||
assert reclaimed.swept is False
|
||||
assert reclaimed.collected is False
|
||||
assert reclaimed.sweep_started_at == 1100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_recovers_stale_claim_without_misreporting_collection(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="stale",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
sweep_started_at=100,
|
||||
),
|
||||
CashuTransaction(
|
||||
token="active",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
sweep_started_at=950,
|
||||
),
|
||||
)
|
||||
receive = AsyncMock(side_effect=RuntimeError("token already spent"))
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", receive),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
receive.assert_awaited_once_with("stale")
|
||||
loaded = await _load(session_factory)
|
||||
assert loaded["stale"].swept is True
|
||||
assert loaded["stale"].collected is False
|
||||
assert loaded["stale"].sweep_started_at is None
|
||||
assert loaded["active"].swept is False
|
||||
assert loaded["active"].sweep_started_at == 950
|
||||
|
||||
@@ -62,6 +62,101 @@ def test_payout_settings_have_sensible_defaults() -> None:
|
||||
assert s.payout_interval_seconds == 900
|
||||
|
||||
|
||||
def test_database_pool_defaults_match_sqlalchemy_capacity() -> None:
|
||||
s = Settings()
|
||||
assert s.database_pool_size == 5
|
||||
assert s.database_max_overflow == 10
|
||||
assert s.database_pool_timeout == 30.0
|
||||
assert s.database_pool_recycle == 1800
|
||||
assert s.database_pool_pre_ping is False
|
||||
assert s.database_pool_hold_warn_seconds == 10.0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "bad_value"),
|
||||
[
|
||||
("database_pool_size", 0),
|
||||
("database_max_overflow", -1),
|
||||
("database_pool_timeout", 0),
|
||||
("database_pool_recycle", -1),
|
||||
("database_pool_hold_warn_seconds", 0),
|
||||
],
|
||||
)
|
||||
def test_database_pool_settings_reject_invalid_values(
|
||||
field: str, bad_value: int
|
||||
) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
Settings.parse_obj({field: bad_value})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_database_pool_fields_are_env_only_not_persisted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""DB pool sizing is infrastructure the node needs *before* it can read the
|
||||
DB, so it can never be configured from the DB — it must never be written to
|
||||
the settings blob, and a stale/injected DB value must never shadow env.
|
||||
"""
|
||||
monkeypatch.setenv("DATABASE_POOL_SIZE", "7")
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
s = await SettingsService.initialize(session)
|
||||
|
||||
# The env value is live for runtime consumers...
|
||||
assert s.database_pool_size == 7
|
||||
# ...but pool sizing is never written to the settings blob.
|
||||
blob = await _read_settings_blob(session)
|
||||
for field in (
|
||||
"database_pool_size",
|
||||
"database_max_overflow",
|
||||
"database_pool_timeout",
|
||||
"database_pool_recycle",
|
||||
"database_pool_pre_ping",
|
||||
"database_pool_hold_warn_seconds",
|
||||
):
|
||||
assert field not in blob
|
||||
|
||||
# Even a stale blob that somehow carries a pool value must not win: env
|
||||
# stays authoritative on the next initialize.
|
||||
await session.exec( # type: ignore
|
||||
text("UPDATE settings SET data = :d WHERE id = 1").bindparams(
|
||||
d=json.dumps({"database_pool_size": 99})
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
again = await SettingsService.initialize(session)
|
||||
assert again.database_pool_size == 7
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_does_not_apply_env_only_fields_to_live_settings(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""DB pool sizing is env-only: a settings update must neither persist it nor
|
||||
mutate the live value. The engine pool is already built at boot from env, so
|
||||
a UI/API update carrying a pool value must not make the live setting diverge
|
||||
from the running pool.
|
||||
"""
|
||||
monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False)
|
||||
monkeypatch.setattr(settings, "database_pool_size", 5)
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
await SettingsService.initialize(session)
|
||||
await SettingsService.update(
|
||||
{"database_pool_size": 99, "name": "PoolTweaker"}, session
|
||||
)
|
||||
|
||||
# A non-env-only field still updates normally...
|
||||
assert settings.name == "PoolTweaker"
|
||||
# ...but the env-only pool size stays at the boot value.
|
||||
assert settings.database_pool_size == 5
|
||||
# ...and it is never written to the settings blob.
|
||||
blob = await _read_settings_blob(session)
|
||||
assert "database_pool_size" not in blob
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"field,bad_value",
|
||||
[
|
||||
|
||||
@@ -71,7 +71,9 @@ async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_sets_reserved_at_on_child_key(session: AsyncSession) -> None:
|
||||
async def test_pay_for_request_sets_reserved_at_on_child_key(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
parent = ApiKey(hashed_key="parentkey", balance=10_000)
|
||||
child = ApiKey(hashed_key="childkey", balance=0, parent_key_hash="parentkey")
|
||||
session.add(parent)
|
||||
@@ -204,7 +206,9 @@ async def test_release_stale_reservations_keeps_fresh(session: AsyncSession) ->
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_stale_reservations_skips_null_reserved_at(session: AsyncSession) -> None:
|
||||
async def test_release_stale_reservations_skips_null_reserved_at(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
# Reservations without a timestamp may belong to instances running older
|
||||
# code (rolling deploy) — the background sweeper must not touch them.
|
||||
key = ApiKey(
|
||||
@@ -224,7 +228,9 @@ async def test_release_stale_reservations_skips_null_reserved_at(session: AsyncS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reset_all_reserved_balances_clears_reserved_at(session: AsyncSession) -> None:
|
||||
async def test_reset_all_reserved_balances_clears_reserved_at(
|
||||
session: AsyncSession,
|
||||
) -> None:
|
||||
key = ApiKey(
|
||||
hashed_key="resetkey",
|
||||
balance=5_000,
|
||||
@@ -404,9 +410,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
||||
AsyncMock(return_value=1_000),
|
||||
),
|
||||
patch.object(proxy_module, "check_token_balance", MagicMock()),
|
||||
patch.object(
|
||||
proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)
|
||||
),
|
||||
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
@@ -418,6 +422,4 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await proxy_module.proxy(request, "v1/chat/completions", session=session)
|
||||
|
||||
revert_mock.assert_awaited_once_with(
|
||||
key, session, 1_000, reservation_snapshot
|
||||
)
|
||||
revert_mock.assert_awaited_once_with(key, session, 950, reservation_snapshot)
|
||||
|
||||
@@ -383,9 +383,7 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
|
||||
AsyncMock(return_value=1_000),
|
||||
),
|
||||
patch.object(proxy_module, "check_token_balance", MagicMock()),
|
||||
patch.object(
|
||||
proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)
|
||||
),
|
||||
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
patch.object(proxy_module, "pay_for_request", AsyncMock(return_value=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
@@ -408,4 +406,4 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
|
||||
assert RAW_ORG_ID not in serialized
|
||||
assert "org-[REDACTED]" in serialized
|
||||
# Single upstream failed -> reservation reverted exactly once (no double-charge).
|
||||
revert_mock.assert_awaited_once_with(key, session, 1_000, reservation)
|
||||
revert_mock.assert_awaited_once_with(key, session, 950, reservation)
|
||||
|
||||
+1077
-41
File diff suppressed because it is too large
Load Diff
@@ -105,6 +105,23 @@ export function DetailedWalletBalance({
|
||||
const formatMintLabel = (detail: BalanceDetail) =>
|
||||
`${detail.mint_url.replace('https://', '').replace('http://', '')} • ${detail.unit.toUpperCase()}`;
|
||||
|
||||
const formatBalanceError = (detail: BalanceDetail) => {
|
||||
const labels: Record<string, string> = {
|
||||
rate_limited: 'rate limited',
|
||||
unreachable: 'unreachable',
|
||||
cooldown: 'cooling down',
|
||||
mint_error: 'mint error',
|
||||
};
|
||||
const label =
|
||||
(detail.error_code ? labels[detail.error_code] : undefined) ??
|
||||
detail.error ??
|
||||
'error';
|
||||
const retryAfter = detail.retry_after_seconds;
|
||||
return retryAfter && retryAfter > 0
|
||||
? `${label} (retry in ${Math.ceil(retryAfter)}s)`
|
||||
: label;
|
||||
};
|
||||
|
||||
return (
|
||||
<>
|
||||
<Card>
|
||||
@@ -262,9 +279,12 @@ export function DetailedWalletBalance({
|
||||
<TableCell className='max-w-md font-mono text-xs break-all whitespace-normal'>
|
||||
{formatMintLabel(detail)}
|
||||
</TableCell>
|
||||
<TableCell className='text-right font-mono'>
|
||||
<TableCell
|
||||
className='text-right font-mono'
|
||||
title={detail.error}
|
||||
>
|
||||
{detail.error
|
||||
? 'error'
|
||||
? formatBalanceError(detail)
|
||||
: formatAmount(walletMsat)}
|
||||
</TableCell>
|
||||
<TableCell className='text-right font-mono'>
|
||||
@@ -306,9 +326,12 @@ export function DetailedWalletBalance({
|
||||
<p className='text-muted-foreground text-xs'>
|
||||
Wallet
|
||||
</p>
|
||||
<p className='font-mono text-sm'>
|
||||
<p
|
||||
className='font-mono text-sm'
|
||||
title={detail.error}
|
||||
>
|
||||
{detail.error
|
||||
? 'error'
|
||||
? formatBalanceError(detail)
|
||||
: formatAmount(walletMsat)}
|
||||
</p>
|
||||
</div>
|
||||
|
||||
@@ -36,6 +36,8 @@ export interface BalanceDetail {
|
||||
user_balance: number;
|
||||
owner_balance: number;
|
||||
error?: string;
|
||||
error_code?: 'rate_limited' | 'unreachable' | 'cooldown' | 'mint_error';
|
||||
retry_after_seconds?: number;
|
||||
}
|
||||
|
||||
export interface WithdrawResponse {
|
||||
|
||||
Reference in New Issue
Block a user