Compare commits

..
18 Commits
Author SHA1 Message Date
9qeklajcandGitHub 3463149b38 Merge pull request #645 from Routstr/fix/increase-default-db-pool
fix(db): increase default connection pool capacity
2026-07-31 02:53:50 +02:00
9qeklajc 19236ecc9d fix(db): increase default connection pool capacity 2026-07-31 02:14:50 +02:00
9qeklajcandGitHub 98fe5a37cd Merge pull request #644 from Routstr/fix/api-key-refund-token-retrieval-main
fix: return persisted API-key refund token
2026-07-31 01:25:18 +02:00
9qeklajc 60566313dc fix: return persisted API-key refund token 2026-07-31 01:19:18 +02:00
9qeklajcandGitHub 88d301398b Merge pull request #642 from Routstr/fix/payout-safety-pool-lifecycle
Fix payout safety and proxy session lifecycle
2026-07-30 03:17:37 +02:00
9qeklajc 4cc9aef61f fix: make payouts and proxy sessions safe 2026-07-30 02:57:49 +02:00
9qeklajcandGitHub c4d27ba02a Merge pull request #634 from Routstr/fix/combined-db-pool-exhaustion
fix: combine DB pool-exhaustion and session-lifecycle fixes
2026-07-30 00:45:00 +02:00
9qeklajc 5ea5024608 resolve review comments 2026-07-29 22:50:33 +02:00
9qeklajc 39f801561b fix: recreate refund sweep migration on latest head 2026-07-26 12:52:24 +02:00
9qeklajc ff55788e2d fix: address PR 634 review feedback 2026-07-26 12:44:37 +02:00
9qeklajc 1131c2d583 test: keep dynamic settings validation mypy-safe 2026-07-24 23:41:01 +02:00
9qeklajc 7108d554c8 merge: combine PR #632 with broader pool-exhaustion fixes 2026-07-24 23:22:01 +02:00
9qeklajc 1eddf89d52 merge: preserve PR #630 history 2026-07-24 23:06:51 +02:00
Jeroen UbbinkandClaude Opus 4.8 a2cedd6769 feat: make the DB connection pool env-configurable
Add DATABASE_POOL_SIZE / DATABASE_MAX_OVERFLOW / DATABASE_POOL_TIMEOUT,
consumed like every other typed env var through the pydantic Settings
(constrained Fields, defaults 5/10/30 matching SQLAlchemy's own baseline
so leaving them unset is behaviour-neutral). An out-of-range or
non-integer value fails validation and refuses to boot — the same
fail-loud behaviour a malformed DATABASE_URL already has — rather than
silently starting up misconfigured. create_db_engine sizes the pool from
these and logs the effective values at startup so they can be confirmed
from the boot output during an incident. In-memory SQLite (StaticPool,
which rejects the pool kwargs) is detected and built without them.

pool_pre_ping is deliberately not exposed: the default backend is a local
SQLite file with no network peer to drop idle connections, so it would add
a SELECT 1 per checkout for no benefit — and it detects dead connections,
not the live-but-wedged ones behind the exhaustion this series addresses.

These knobs are infrastructure the node needs before it can open a DB
session, so they can never be sourced from the DB (chicken-and-egg). A new
ENV_ONLY_FIELDS set keeps them out of the persisted settings blob and
stops a DB value from shadowing env in both SettingsService.initialize and
.update, so env stays authoritative.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-24 22:26:57 +02:00
Jeroen UbbinkandClaude Opus 4.8 7b0ade3987 fix: stop AsyncSession pool exhaustion in balance reads
fetch_all_balances shared a single AsyncSession across the tasks it ran
with asyncio.gather. AsyncSession is not safe for concurrent use: the
overlapping queries raised "concurrent operations are not permitted" and
left connections wedged until the QueuePool was exhausted, after which
admin/API endpoints returned 500/502 until a container restart.

Read all outstanding user liabilities up front in one short-lived session
with a single grouped query (balances_by_mint_and_unit), then run the
per-mint balance checks concurrently with no session in scope. A failure
reading liabilities now degrades gracefully — the page still reports the
known wallet custody and blanks only the unknowable user/owner split,
tagging each mint with the error — instead of 500-ing the whole page.

periodic_payout reads each liability fresh, immediately before the payout
decision, rather than from a single pre-loop snapshot: the per-mint round
trip is slow, and a user top-up during the cycle would otherwise let a
later mint/unit act on a stale-low liability and over-send funds owed to
users (related to the payout-safety concern in issue #611).

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-24 22:26:25 +02:00
9qeklajc 2410a4a6ce fix: address review findings on payout liability staleness, sweep races, and cancellation safety
- periodic_payout: fetch liability per mint/unit right before computing
  available balance, so a concurrent top-up can only shrink the payout
- refund sweep: atomically claim each refund before redeeming; release the
  claim on failure so retries still happen and concurrent sweeps cannot
  misreport a sweep as client-collected
- check_invoice_payment: catch BaseException so task cancellation after a
  successful mint still emits the reconciliation alert
- tests: DB-guard race test where both mints succeed (exactly one credit);
  pool_size=1 test proving the fee payout releases its connection during
  the external send
2026-07-24 22:14:57 +02:00
9qeklajc 1b09639265 docs: document DB pool and mint concurrency env vars in .env.example 2026-07-24 21:44:11 +02:00
9qeklajc 7394b10e75 fix-concurent-issue 2026-07-24 21:29:21 +02:00
32 changed files with 2717 additions and 2861 deletions
+15
View File
@@ -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=10
# DATABASE_MAX_OVERFLOW=20
# DATABASE_POOL_TIMEOUT=15
# 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
-56
View File
@@ -327,62 +327,6 @@ GET /v1/models
}
```
### List Model Paths
Get the selectable upstream routes for each advertised model. This endpoint is
discovery-only; request-side selection will be added separately.
```http
GET /v1/models/paths
```
**Response:**
```json
{
"data": [
{
"id": "anthropic/claude-sonnet-4",
"paths": [
{
"path": "url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=anthropic%2Fclaude-sonnet-4",
"provider": {"id": 12, "slug": "anthropic-primary", "type": "anthropic"},
"endpoint": null
},
{
"path": "url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=anthropic%2Fclaude-sonnet-4&endpoint=google-vertex%2Fus",
"provider": {"id": 42, "slug": "openrouter-main", "type": "openrouter"},
"endpoint": {"tag": "google-vertex/us", "name": "Google"}
}
]
}
],
"updated_at": 1753500000
}
```
`path` is an opaque, percent-encoded selector. Clients must store and return it
unchanged rather than parsing or reconstructing it. It identifies the exact
configured route with `url`, `provider-id`, and `model-id`. To avoid exposing
private network details, a configured private IP address or any URL with an
explicit port is advertised as `http://localhost`. OpenRouter routes additionally
preserve the exact machine-readable endpoint `tag`. Provider slugs/types and
endpoint names remain display data. When request-side selection is implemented,
an endpoint tag must not silently fall back to another backend.
### List Paths for One Model
Use the exact model ID advertised by `/v1/models`. The query parameter safely
supports IDs containing `/`.
```http
GET /v1/models/paths/model?model_id=anthropic/claude-sonnet-4
```
The response uses the same path objects and `updated_at` field as the collection
endpoint. An unknown model returns `404 Model not found`. A known model whose
paths have not been discovered yet returns `200` with an empty `data` array.
## Wallet Management
### Create Wallet (Coming Soon)
+1 -2
View File
@@ -100,7 +100,6 @@ All errors follow a consistent format:
Standard OpenAI-compatible endpoints:
- **Models**: `/v1/models`
- **Model paths**: `/v1/models/paths`, `/v1/models/paths/model?model_id=...`
- **Responses**: `/v1/responses`
- **Chat Completions**: `/v1/chat/completions`
- **Embeddings**: `/v1/embeddings`
@@ -303,7 +302,7 @@ Get node metadata:
GET /v1/info
```
Supported models and pricing are available at `/v1/models`. Upstream provider path discovery is available at `/v1/models/paths` and `/v1/models/paths/model?model_id=...`.
Supported models and pricing are available at `/v1/models`.
## Next Steps
-7
View File
@@ -142,8 +142,6 @@ Use environment variables for:
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
| `CORS_ORIGINS` | Allowed CORS origins | `*` |
| `RELAYS` | Nostr relays (comma-separated) | (default set) |
| `MODEL_PATHS_REFRESH_INTERVAL_SECONDS` | How often to refresh `/v1/models/paths` discovery data; set `0` to pause the refresh (previously discovered paths keep being served) | `600` |
| `ENABLE_MODEL_PATHS_REFRESH` | Kill switch for the background model-path refresh (OpenRouter endpoint fan-out) | `true` |
### Priority
@@ -177,8 +175,3 @@ Manage which AI models you offer:
- **Create aliases** — friendly names for models
See [Pricing](pricing.md) for per-model pricing strategies.
Model path discovery is refreshed in the background and exposed through
`/v1/models/paths`. The response groups each client-visible model ID with the
provider paths that may appear in chat-completion response metadata. Tune the
refresh cadence with `MODEL_PATHS_REFRESH_INTERVAL_SECONDS`.
@@ -1,54 +0,0 @@
"""add model paths table
Revision ID: 4e0c3d195a49
Revises: 9c4d8e2f1a6b
Create Date: 2026-07-24 21:14:39.062179
"""
import sqlalchemy as sa
import sqlmodel
from alembic import op
# revision identifiers, used by Alembic.
revision = "4e0c3d195a49"
down_revision = "9c4d8e2f1a6b"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"model_paths",
sa.Column("id", sa.Integer(), nullable=False),
sa.Column("model_id", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("path", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("provider_slug", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("provider_type", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
sa.Column("endpoint_tag", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
sa.Column("endpoint_name", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
sa.Column("upstream_provider_id", sa.Integer(), nullable=False),
sa.Column("updated_at", sa.Integer(), nullable=False, server_default="0"),
sa.ForeignKeyConstraint(
["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE"
),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint(
"model_id",
"path",
"upstream_provider_id",
name="uq_model_paths_model_path_provider",
),
)
# No standalone index on model_id: the unique constraint's autoindex already
# leads on model_id.
op.create_index(
op.f("ix_model_paths_upstream_provider_id"),
"model_paths",
["upstream_provider_id"],
unique=False,
)
def downgrade() -> None:
op.drop_index(op.f("ix_model_paths_upstream_provider_id"), table_name="model_paths")
op.drop_table("model_paths")
@@ -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")
+24
View File
@@ -33,6 +33,7 @@ from .wallet import (
classify_redemption_error,
credit_balance,
deserialize_token_from_string,
wallet_operation_guard,
)
if TYPE_CHECKING:
@@ -153,6 +154,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.
+32
View File
@@ -259,6 +259,36 @@ async def _lookup_key_no_create(
return None
async def _get_persisted_api_key_refund(
key: ApiKey, session: AsyncSession
) -> dict[str, str] | None:
result = await session.exec(
select(CashuTransaction)
.where(
CashuTransaction.api_key_hashed_key == key.hashed_key,
CashuTransaction.type == "out",
CashuTransaction.source == "apikey",
)
.order_by(col(CashuTransaction.created_at).desc())
)
refund = result.first()
if refund is None:
return None
if refund.swept:
raise HTTPException(status_code=410, detail="Refund has been swept")
refund.collected = True
session.add(refund)
await session.commit()
persisted = {"token": refund.token}
if refund.unit == "sat":
persisted["sats"] = str(refund.amount)
else:
persisted["msats"] = str(refund.amount)
return persisted
async def _restore_balance(
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
) -> None:
@@ -353,6 +383,8 @@ async def refund_wallet_endpoint(
if key.total_balance <= 0:
if cached := await _refund_cache_get(bearer_value):
return cached
if persisted := await _get_persisted_api_key_refund(key, session):
return persisted
if key.parent_key_hash:
raise HTTPException(
+8 -16
View File
@@ -51,13 +51,6 @@ ADMIN_SESSION_DURATION = 3600
MAX_USAGE_ANALYTICS_HOURS = 365 * 24
async def _refresh_provider_model_paths(upstream_provider_id: int) -> None:
"""Queue discovery sync without blocking the committed admin mutation."""
from ..upstream.model_paths import schedule_model_paths_refresh_for_provider
await schedule_model_paths_refresh_for_provider(upstream_provider_id)
async def require_admin_api(request: Request) -> None:
auth_header = request.headers.get("Authorization")
if not auth_header or not auth_header.startswith("Bearer "):
@@ -586,7 +579,6 @@ async def upsert_provider_model(
await session.refresh(row)
await refresh_model_maps()
await _refresh_provider_model_paths(provider_pk)
return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee
).dict() # type: ignore
@@ -641,7 +633,6 @@ async def delete_provider_model(provider_id: str, model_id: str) -> dict[str, ob
await session.delete(row)
await session.commit()
await refresh_model_maps()
await _refresh_provider_model_paths(provider_pk)
return {"ok": True, "deleted_id": model_id}
@@ -661,7 +652,6 @@ async def delete_all_provider_models(provider_id: str) -> dict[str, object]:
await session.delete(row) # type: ignore
await session.commit()
await refresh_model_maps()
await _refresh_provider_model_paths(provider_pk)
return {"ok": True, "deleted": len(rows)}
@@ -753,7 +743,6 @@ async def batch_override_provider_models(
await session.commit()
await refresh_model_maps()
await _refresh_provider_model_paths(provider_pk)
return {
"ok": True,
"count": overridden_count,
@@ -954,7 +943,6 @@ async def create_upstream_provider(
await reinitialize_upstreams()
await refresh_model_maps()
await _refresh_provider_model_paths(_provider_pk(provider))
return _serialize_provider(provider)
@@ -980,7 +968,6 @@ async def update_upstream_provider(
await reinitialize_upstreams()
await refresh_model_maps()
await _refresh_provider_model_paths(_provider_pk(provider))
return _serialize_provider(provider)
@@ -1016,7 +1003,6 @@ async def update_upstream_provider_by_slug(
await reinitialize_upstreams()
await refresh_model_maps()
await _refresh_provider_model_paths(_provider_pk(provider))
return _serialize_provider(provider)
@@ -1683,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,
}
+115 -64
View File
@@ -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
@@ -258,7 +317,9 @@ async def prune_dead_api_keys(session: AsyncSession, min_age_seconds: int) -> in
.where(col(ApiKey.total_spent) == 0)
.where(col(ApiKey.total_requests) == 0)
.where(col(ApiKey.parent_key_hash).is_(None))
.where((col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff))
.where(
(col(ApiKey.created_at).is_(None)) | (col(ApiKey.created_at) < cutoff)
)
.where(~pending_invoice)
.where(~has_children)
)
@@ -312,60 +373,6 @@ class ModelRow(SQLModel, table=True): # type: ignore
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
class ModelPathRow(SQLModel, table=True): # type: ignore
"""Upstream provider path a model is reachable through.
Discovery/visibility data only. ``model_id`` is intentionally NOT globally
unique: it is the client-visible ``/v1/models`` id (``forwarded_model_id or
id``) grouped across every provider that exposes the model. A single model
can therefore have several rows — one per direct provider path plus one per
OpenRouter sub-provider endpoint.
"""
__tablename__ = "model_paths"
__table_args__ = (
UniqueConstraint(
"model_id",
"path",
"upstream_provider_id",
name="uq_model_paths_model_path_provider",
),
)
id: int | None = Field(default=None, primary_key=True)
# No standalone index on model_id: the unique constraint's autoindex already
# leads on model_id, so a second index only adds write amplification.
model_id: str = Field(
description="Client-visible /v1/models id (forwarded_model_id or id)"
)
path: str = Field(
description=(
"Opaque selector containing upstream URL, provider ID, model ID, "
"and optional endpoint tag"
)
)
provider_slug: str = Field(
description="Public slug of the configured upstream provider"
)
provider_type: str = Field(description="Configured upstream provider type")
endpoint_tag: str | None = Field(
default=None,
description="Exact OpenRouter endpoint tag used for request-side selection",
)
endpoint_name: str | None = Field(
default=None, description="Human-readable endpoint display name"
)
upstream_provider_id: int = Field(
index=True,
foreign_key="upstream_providers.id",
ondelete="CASCADE",
description="upstream_providers.id this path was discovered from",
)
updated_at: int = Field(
default=0,
description="Unix timestamp of the refresh cycle that wrote this row",
)
class LightningInvoice(SQLModel, table=True): # type: ignore
__tablename__ = "lightning_invoices"
@@ -375,7 +382,8 @@ 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"
@@ -420,6 +428,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",
@@ -651,7 +663,9 @@ class CliToken(SQLModel, table=True): # type: ignore
"""Long-lived authorization token for CLI/agent use against admin endpoints."""
__tablename__ = "cli_tokens"
id: str = Field(primary_key=True, default_factory=lambda: uuid.uuid4().hex)
id: str = Field(
primary_key=True, default_factory=lambda: uuid.uuid4().hex
)
token: str = Field(unique=True, index=True, description="Bearer token value")
name: str = Field(description="Human-readable label for this token")
created_at: int = Field(default_factory=lambda: int(time.time()))
@@ -750,7 +764,9 @@ async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> bool:
return result.rowcount == 1
async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) -> bool:
async def complete_routstr_fee_payout(
session: AsyncSession, paid_msats: int
) -> bool:
"""Mark a checkpointed payout complete after the external payment succeeds."""
stmt = (
update(RoutstrFee)
@@ -768,14 +784,49 @@ async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) ->
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:
+9 -15
View File
@@ -58,7 +58,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
providers_task = None
models_refresh_task = None
model_maps_refresh_task = None
model_paths_refresh_task = None
key_reset_task = None
stale_reservation_task = None
dead_key_prune_task = None
@@ -131,13 +130,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
refresh_upstreams_models_periodically(get_upstreams)
)
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
# Always started: the loop re-reads the enable flag and interval every
# iteration, so 0 -> N (or re-enabling) takes effect without a restart.
from ..upstream.model_paths import refresh_model_paths_periodically
model_paths_refresh_task = asyncio.create_task(
refresh_model_paths_periodically(get_upstreams)
)
payout_task = asyncio.create_task(periodic_payout())
if global_settings.nsec:
nip91_task = asyncio.create_task(announce_provider())
@@ -145,7 +137,9 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
if global_settings.providers_refresh_interval_seconds > 0:
providers_task = asyncio.create_task(providers_cache_refresher())
key_reset_task = asyncio.create_task(periodic_key_reset())
stale_reservation_task = asyncio.create_task(periodic_stale_reservation_sweep())
stale_reservation_task = asyncio.create_task(
periodic_stale_reservation_sweep()
)
dead_key_prune_task = asyncio.create_task(periodic_dead_key_prune())
auto_topup_task = asyncio.create_task(periodic_auto_topup())
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
@@ -182,8 +176,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
models_refresh_task.cancel()
if model_maps_refresh_task is not None:
model_maps_refresh_task.cancel()
if model_paths_refresh_task is not None:
model_paths_refresh_task.cancel()
if key_reset_task is not None:
key_reset_task.cancel()
if stale_reservation_task is not None:
@@ -217,8 +209,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
tasks_to_wait.append(models_refresh_task)
if model_maps_refresh_task is not None:
tasks_to_wait.append(model_maps_refresh_task)
if model_paths_refresh_task is not None:
tasks_to_wait.append(model_paths_refresh_task)
if key_reset_task is not None:
tasks_to_wait.append(key_reset_task)
if stale_reservation_task is not None:
@@ -255,7 +245,9 @@ class _ImmutableStaticFiles(StaticFiles):
async def get_response(self, path: str, scope: Scope) -> StarletteResponse:
response = await super().get_response(path, scope)
if response.status_code == 200:
response.headers["Cache-Control"] = "public, max-age=31536000, immutable"
response.headers["Cache-Control"] = (
"public, max-age=31536000, immutable"
)
return response
@@ -329,7 +321,9 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
# Serve the App Router RSC payload for the home page.
@app.get("/index.txt", include_in_schema=False)
async def serve_root_rsc() -> FileResponse:
return FileResponse(UI_DIST_PATH / "index.txt", media_type="text/x-component")
return FileResponse(
UI_DIST_PATH / "index.txt", media_type="text/x-component"
)
# Next.js is built with `trailingSlash: true`, so all UI page URLs end
# with a slash (e.g. `/login/`). The proxy router catches `/{path:path}`
+57 -14
View File
@@ -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
@@ -95,17 +98,27 @@ class Settings(BaseSettings):
models_refresh_interval_seconds: int = Field(
default=360, env="MODELS_REFRESH_INTERVAL_SECONDS"
)
model_paths_refresh_interval_seconds: int = Field(
default=600, env="MODEL_PATHS_REFRESH_INTERVAL_SECONDS"
)
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
enable_model_paths_refresh: bool = Field(
default=True, env="ENABLE_MODEL_PATHS_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 provide
# headroom for Routstr's concurrent request and background-payment workload.
# 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=10, ge=1, env="DATABASE_POOL_SIZE")
database_max_overflow: int = Field(default=20, ge=0, env="DATABASE_MAX_OVERFLOW")
database_pool_timeout: float = Field(
default=15.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
@@ -125,8 +138,9 @@ 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."""
@@ -151,10 +165,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:
@@ -337,7 +373,9 @@ class SettingsService:
{
k: v
for k, v in db_json.items()
if v not in (None, "", [], {}) and k in valid_fields
if v not in (None, "", [], {})
and k in valid_fields
and k not in ENV_ONLY_FIELDS
}
)
merged_dict = Settings(**merged_dict).dict()
@@ -405,8 +443,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
+143 -42
View File
@@ -5,13 +5,14 @@ import time
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 get_wallet, wallet_operation_guard
logger = get_logger(__name__)
@@ -222,80 +223,180 @@ async def recover_invoice(
async def check_invoice_payment(
invoice: LightningInvoice, session: AsyncSession
) -> None:
# Minting makes proofs visible before database finalization. Share the
# cross-process wallet guard with owner payout so that visibility and the
# corresponding liability commit are observed atomically by the payout loop.
async with wallet_operation_guard():
await _check_invoice_payment_locked(invoice, session)
async def _check_invoice_payment_locked(
invoice: LightningInvoice, session: AsyncSession
) -> None:
minted = False
invoice_id = invoice.id
invoice_purpose = invoice.purpose
invoice_amount_sats = invoice.amount_sats
invoice_payment_hash = invoice.payment_hash
finalized_api_key_hash = invoice.api_key_hash
try:
# A preceding invoice lookup starts a transaction. End it before the
# potentially slow mint request so it cannot pin a pool connection.
await session.commit()
wallet = await get_wallet(settings.primary_mint, "sat")
mint_status = await wallet.get_mint_quote(invoice_payment_hash)
if not mint_status.paid:
return
mint_status = await wallet.get_mint_quote(invoice.payment_hash)
# Do not redeem a paid top-up quote if its target has already been
# pruned. This validation owns a short-lived session and releases its
# connection before mint redemption starts.
if invoice_purpose == "topup":
if not finalized_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, finalized_api_key_hash)
if target is None:
terminal = await validation_session.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(
col(LightningInvoice.id) == invoice_id,
col(LightningInvoice.status) == "pending",
)
.values(status="reconciliation_required")
)
await validation_session.commit()
if terminal.rowcount == 1:
set_committed_value(
invoice, "status", "reconciliation_required"
)
else:
committed_invoice = await validation_session.get(
LightningInvoice, invoice_id
)
if committed_invoice is not None:
set_committed_value(
invoice, "status", committed_invoice.status
)
logger.critical(
"Paid topup invoice target API key was not found; reconciliation required",
extra={"invoice_id": invoice_id},
)
return
if mint_status.paid:
invoice.status = "paid"
invoice.paid_at = int(time.time())
# The mint enforces single-use quotes, so a concurrent checker that
# races us here fails inside wallet.mint rather than double-minting.
await wallet.mint(invoice_amount_sats, quote_id=invoice_payment_hash)
minted = True
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)
# Paid finalization owns a fresh session. The API/watcher session is
# never rolled back by this function, so its invoice and sibling ORM
# objects remain usable after a DB failure or lost CAS race.
async with create_session() as finalization_session:
if invoice_purpose == "create":
api_key = await _create_api_key_record(invoice, finalization_session)
finalized_api_key_hash = api_key.hashed_key
elif invoice_purpose == "topup":
await _credit_topup_record(invoice, finalization_session)
await session.commit()
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
else None,
},
# Conditional transition guards against double-credit: the credit
# above and this status flip commit atomically, and a lost race
# rolls both back in the owned finalization session.
paid_at = int(time.time())
finalized = await finalization_session.exec( # type: ignore[call-overload]
update(LightningInvoice)
.where(
col(LightningInvoice.id) == invoice_id,
col(LightningInvoice.status) == "pending",
)
.values(
status="paid",
paid_at=paid_at,
api_key_hash=finalized_api_key_hash,
)
)
if finalized.rowcount != 1:
await finalization_session.rollback()
committed_invoice = await finalization_session.get(
LightningInvoice, invoice_id
)
await finalization_session.commit()
if committed_invoice is not None:
# A concurrent finalizer won the CAS. Publish only the
# state observed from the database after ending the owned
# read transaction; never refresh the caller's session.
set_committed_value(
invoice, "api_key_hash", committed_invoice.api_key_hash
)
set_committed_value(invoice, "status", committed_invoice.status)
set_committed_value(invoice, "paid_at", committed_invoice.paid_at)
return
await finalization_session.commit()
except Exception as e:
# Only publish finalized values to the caller-owned object after the
# owned transaction has committed successfully.
set_committed_value(invoice, "api_key_hash", finalized_api_key_hash)
set_committed_value(invoice, "status", "paid")
set_committed_value(invoice, "paid_at", paid_at)
logger.info(
"Lightning invoice paid",
extra={
"invoice_id": invoice_id,
"amount_sats": invoice_amount_sats,
"purpose": invoice_purpose,
"api_key_hash": finalized_api_key_hash[:8] + "..."
if finalized_api_key_hash
else None,
},
)
except BaseException as e:
# BaseException so task cancellation (e.g. client disconnect) after a
# successful mint still triggers the reconciliation alert. Any rollback
# belongs to create_session(), never to the caller-owned session.
if minted:
logger.critical(
"Invoice mint succeeded but DB finalization failed; reconciliation required",
extra={"invoice_id": invoice_id, "purpose": invoice_purpose},
)
if not isinstance(e, Exception):
raise
logger.error(f"Failed to check invoice payment: {e}")
async def create_api_key_from_invoice(
async def _create_api_key_record(
invoice: LightningInvoice, session: AsyncSession
) -> ApiKey:
wallet = await get_wallet(settings.primary_mint, "sat")
await wallet.mint(invoice.amount_sats, quote_id=invoice.payment_hash)
dummy_token = f"invoice-{invoice.id}-{invoice.payment_hash}"
hashed_key = hashlib.sha256(dummy_token.encode()).hexdigest()
api_key = ApiKey(
hashed_key=hashed_key,
balance=invoice.amount_sats * 1000, # Convert to msats
balance=invoice.amount_sats * 1000,
refund_currency="sat",
refund_mint_url=settings.primary_mint,
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(
async def _credit_topup_record(
invoice: LightningInvoice, 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:
credited = 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)
)
if credited.rowcount != 1:
raise ValueError("Associated API key not found")
api_key.balance += invoice.amount_sats * 1000 # Convert to msats
await session.flush()
INVOICE_WATCH_INTERVAL_SECONDS = 5
INVOICE_WATCH_BATCH_LIMIT = 100
+4 -35
View File
@@ -455,9 +455,7 @@ async def _update_sats_pricing_once() -> None:
for m in upstream.get_cached_models()
]
upstream._models_cache = updated_models
upstream._models_by_id = {
m.forwarded_model_id or m.id: m for m in updated_models
}
upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models}
updated_count += len(updated_models)
if updated_count > 0:
@@ -512,7 +510,9 @@ class ModelTestRequest(V2BaseModel):
request_data: dict
@models_router.post("/api/models/test", dependencies=[Depends(_require_admin_api)])
@models_router.post(
"/api/models/test", dependencies=[Depends(_require_admin_api)]
)
async def test_model(
payload: ModelTestRequest,
session: AsyncSession = Depends(get_session),
@@ -595,37 +595,6 @@ async def test_model(
}
@models_router.get("/v1/models/paths")
@models_router.get("/v1/models/paths/", include_in_schema=False)
async def model_paths() -> dict:
"""All models with every upstream provider path they are reachable through."""
from ..upstream.model_paths import get_all_model_paths
return await get_all_model_paths()
@models_router.get("/v1/models/paths/model")
@models_router.get("/v1/models/paths/model/", include_in_schema=False)
async def model_paths_for_model(model_id: str) -> dict:
"""Paths for a single model.
Uses a query parameter (``?model_id=...``) under a fully static route so
model ids containing ``/`` (e.g. ``anthropic/claude-opus-4.6``) need no URL
encoding and there is no dynamic-route ambiguity.
"""
from ..proxy import get_unique_models
from ..upstream.model_paths import get_paths_for_model
result = await get_paths_for_model(model_id)
if not result["data"]:
advertised_ids = {
model.forwarded_model_id or model.id for model in get_unique_models()
}
if model_id not in advertised_ids:
raise HTTPException(status_code=404, detail="Model not found")
return result
@models_router.get("/v1/models")
@models_router.get("/v1/models/", include_in_schema=False)
@models_router.get("/models")
+27 -13
View File
@@ -1,4 +1,5 @@
import asyncio
import inspect
import json
from typing import Any
@@ -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
@@ -183,19 +191,6 @@ async def refresh_model_maps() -> None:
disabled_model_keys=disabled_model_keys,
)
# Keep model-path discovery in sync with admin mutations: disabling or
# deleting a provider must stop advertising its paths immediately rather
# than after the next timed refresh.
from .upstream.model_paths import prune_model_paths_for_inactive_providers
try:
await prune_model_paths_for_inactive_providers()
except Exception as e: # noqa: BLE001 - discovery sync must not break routing
logger.warning(
"Failed to prune model paths for inactive providers",
extra={"error": str(e), "error_type": type(e).__name__},
)
async def refresh_model_maps_periodically() -> None:
"""Background task to refresh model maps every minute."""
@@ -233,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
@@ -463,6 +472,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
@@ -494,8 +506,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))
-807
View File
@@ -1,807 +0,0 @@
"""Model-path discovery service.
Exposes every selectable upstream route a Routstr model is reachable through.
This PR remains discovery-only: request-side routing will consume the opaque
selectors in a follow-up.
A path is a standard percent-encoded query string containing the configured
upstream URL, provider ID, client-visible model ID and, for an exact OpenRouter
endpoint, its machine-readable tag::
url=https%3A%2F%2Fapi.anthropic.com%2Fv1&provider-id=12&model-id=claude-sonnet-4
url=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1&provider-id=42&model-id=claude-sonnet-4&endpoint=google-vertex%2Fus
"""
from __future__ import annotations
import asyncio
import ipaddress
import random
import time
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Callable
from urllib.parse import urlencode, urlsplit
import httpx
from sqlalchemy.dialects.sqlite import insert
from sqlalchemy.orm import selectinload
from sqlmodel import col, delete, select
from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_session
from ..core.logging import get_logger
if TYPE_CHECKING:
from .base import BaseUpstreamProvider
logger = get_logger(__name__)
# Bound the per-model OpenRouter /endpoints fan-out so a provider with hundreds
# of models does not open hundreds of concurrent requests every refresh.
_OPENROUTER_CONCURRENCY = 5
_OPENROUTER_TIMEOUT_SECONDS = 10.0
# Rows inserted per statement during persist. Keeps each INSERT bounded while
# avoiding per-row round-trips that hold SQLite's write lock for ~1s per cycle.
_PERSIST_CHUNK_SIZE = 500
# Admin mutations enqueue provider IDs here instead of running OpenRouter's
# per-model endpoint fan-out inside the request. One worker serializes refreshes
# and coalesces repeated mutations for the same provider.
_scheduled_provider_refresh_ids: set[int] = set()
_scheduled_provider_refresh_task: asyncio.Task[None] | None = None
# Visibility key used across this module: routing carries the provider
# dimension everywhere (ModelRow's primary key is (id, upstream_provider_id)),
# so all model-id keyed maps here do too, lowercased like proxy.refresh_model_maps.
ModelKey = tuple[str, int]
@dataclass(frozen=True)
class EndpointIdentity:
"""Exact OpenRouter endpoint identity returned by ``/endpoints``."""
tag: str
provider_name: str | None
@dataclass(frozen=True)
class ConfiguredProviderIdentity:
"""Public identity of one configured upstream provider."""
id: int
slug: str
provider_type: str
base_url: str
@dataclass(frozen=True)
class DiscoveredPath:
"""One model route ready for persistence and API serialization."""
model_id: str
path: str
provider: ConfiguredProviderIdentity
endpoint_tag: str | None = None
endpoint_name: str | None = None
@dataclass(frozen=True)
class ProviderPathSnapshot:
"""Refresh result plus model IDs whose prior rows must survive degradation."""
paths: tuple[DiscoveredPath, ...]
preserve_model_ids: frozenset[str] = frozenset()
def public_provider_url(base_url: str) -> str:
"""Mask private IP addresses and URLs with explicit ports."""
parsed = urlsplit(base_url)
try:
if parsed.port is not None:
return "http://localhost"
except ValueError:
# An invalid explicit port must not accidentally leak through.
return "http://localhost"
hostname = parsed.hostname
if hostname is None:
return base_url
try:
address = ipaddress.ip_address(hostname)
except ValueError:
return base_url
return "http://localhost" if address.is_private else base_url
def encode_model_path(
base_url: str,
provider_id: int,
model_id: str,
endpoint_tag: str | None = None,
) -> str:
"""Encode the complete upstream route selector advertised to clients."""
components: list[tuple[str, str | int]] = [
("url", base_url),
("provider-id", provider_id),
("model-id", model_id),
]
if endpoint_tag:
components.append(("endpoint", endpoint_tag))
return urlencode(components)
def _make_http_client() -> httpx.AsyncClient:
"""Client factory, separated so tests can substitute a mock transport."""
return httpx.AsyncClient()
def is_openrouter_base_url(base_url: str | None) -> bool:
"""True when ``base_url`` points at OpenRouter.
Deliberately separate from ``BaseUpstreamProvider._upstream_accepts_cache_control``:
that predicate also returns True for native Anthropic (correct for
cache-control, wrong for OpenRouter endpoint discovery). This one keys only
on the URL so a ``GenericUpstreamProvider`` aimed at OpenRouter is matched
while native Anthropic is not.
"""
return "openrouter.ai" in (base_url or "")
def exposed_model_id(model: object) -> str:
"""Return exactly the ID advertised by ``/v1/models``.
A forwarded ID is already a public routable alias and must remain intact,
including any slash. Without one, ``/v1/models`` exposes the base ID.
"""
forwarded = getattr(model, "forwarded_model_id", None)
if forwarded:
return forwarded
return public_model_id(getattr(model, "id"))
def public_model_id(model_id: str) -> str:
"""Model id exposed by model-path API responses.
Uses the same rule as ``create_model_mappings.get_base_model_id`` and
``resolve_model_alias`` — strip everything before the *first* slash — so
the id shown here can be sent back to ``/v1/chat/completions`` verbatim.
"""
return model_id.split("/", 1)[1] if "/" in model_id else model_id
def openrouter_author_slug(model: object) -> str | None:
"""Return a canonical ``author/slug`` for the OpenRouter endpoints API.
Prefer ``canonical_slug``, then a slash-containing ``id``, then a
slash-containing ``forwarded_model_id``. The forwarded id is exactly what
the proxy sends upstream for admin-created alias rows (``base.py`` forwards
``forwarded_model_id or id``), so it is a valid OpenRouter id when the
bare ``id`` is a local alias with no slash.
"""
canonical = getattr(model, "canonical_slug", None)
if canonical and "/" in canonical:
return canonical
model_id = getattr(model, "id", None)
if model_id and "/" in model_id:
return model_id
forwarded = getattr(model, "forwarded_model_id", None)
if forwarded and "/" in forwarded:
return forwarded
return None
class _RefreshCycleState:
"""Per-refresh shared state: fetch dedupe cache and rate-limit latch.
``endpoint_cache`` dedupes byte-identical ``/endpoints`` fetches when two
providers point at the same OpenRouter base URL. ``rate_limited`` latches
on the first 429 so the rest of the cycle stops hammering a throttled API;
the whole provider result then degrades to "unknown" instead of an empty
list, which preserves previously persisted rows.
"""
def __init__(self) -> None:
self.endpoint_cache: dict[tuple[str, str], list[EndpointIdentity] | None] = {}
self.rate_limited = False
async def _fetch_openrouter_endpoint_subproviders(
client: httpx.AsyncClient,
base_url: str,
api_key: str,
author_slug: str,
semaphore: asyncio.Semaphore,
cycle: _RefreshCycleState,
) -> list[EndpointIdentity] | None:
"""Return exact endpoint identities for one model, or ``None`` when unknown.
``None`` (not ``[]``) signals a degraded fetch — network failure, rate
limit, non-200, or an unparseable payload — so callers can distinguish
"this model has no endpoints" from "we could not find out". Failures are
logged and swallowed so one model never breaks the whole refresh.
"""
cache_key = (base_url, author_slug)
if cache_key in cycle.endpoint_cache:
return cycle.endpoint_cache[cache_key]
if cycle.rate_limited:
return None
url = f"{base_url.rstrip('/')}/models/{author_slug}/endpoints"
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
result: list[EndpointIdentity] | None
async with semaphore:
try:
resp = await client.get(
url, headers=headers, timeout=_OPENROUTER_TIMEOUT_SECONDS
)
except Exception as e: # noqa: BLE001 - isolate per-model failures
logger.warning(
"OpenRouter endpoint discovery request failed",
extra={"author_slug": author_slug, "error": str(e)},
)
cycle.endpoint_cache[cache_key] = None
return None
if resp.status_code == 429:
logger.warning(
"OpenRouter endpoint discovery rate-limited; aborting cycle",
extra={"author_slug": author_slug},
)
cycle.rate_limited = True
cycle.endpoint_cache[cache_key] = None
return None
if resp.status_code != 200:
logger.warning(
"OpenRouter endpoint discovery non-200",
extra={"author_slug": author_slug, "status_code": resp.status_code},
)
cycle.endpoint_cache[cache_key] = None
return None
try:
payload = resp.json()
data = payload.get("data") if isinstance(payload, dict) else None
endpoints = data.get("endpoints") if isinstance(data, dict) else None
if not isinstance(endpoints, list):
raise ValueError("endpoints must be a list")
identities: dict[str, EndpointIdentity] = {}
for endpoint in endpoints:
if not isinstance(endpoint, dict):
continue
tag = endpoint.get("tag")
if not isinstance(tag, str) or not tag.strip():
continue
provider_name = endpoint.get("provider_name")
identities.setdefault(
tag,
EndpointIdentity(
tag=tag,
provider_name=provider_name
if isinstance(provider_name, str) and provider_name
else None,
),
)
if endpoints and not identities:
raise ValueError("endpoints contain no usable tags")
result = list(identities.values())
except Exception as e: # noqa: BLE001
logger.warning(
"OpenRouter endpoint discovery bad payload",
extra={"author_slug": author_slug, "error": str(e)},
)
result = None
cycle.endpoint_cache[cache_key] = result
return result
async def _load_model_visibility() -> tuple[
dict[ModelKey, ModelRow],
set[ModelKey],
dict[int, ConfiguredProviderIdentity],
]:
"""Load the same DB model visibility inputs used by routing.
``refresh_model_maps`` builds routing from enabled providers, enabled DB
override rows, and disabled model keys — all keyed on
``(model_id.lower(), upstream_provider_id)`` because ``ModelRow``'s primary
key is composite and the same id legitimately exists on several providers.
Model-path discovery uses the same keying so disabling a model on one
provider never hides it on another, and one provider's
``forwarded_model_id`` alias is never applied to a different provider.
"""
async with create_session() as session:
query = select(UpstreamProviderRow).options(
selectinload(UpstreamProviderRow.models) # type: ignore[arg-type]
)
provider_rows = (await session.exec(query)).all()
overrides_by_key: dict[ModelKey, ModelRow] = {}
disabled_model_keys: set[ModelKey] = set()
provider_identities: dict[int, ConfiguredProviderIdentity] = {}
for provider in provider_rows:
if not provider.enabled or provider.id is None:
continue
provider_identities[provider.id] = ConfiguredProviderIdentity(
id=provider.id,
slug=provider.slug or f"provider-{provider.id}",
provider_type=provider.provider_type,
base_url=public_provider_url(provider.base_url),
)
for model in provider.models:
key = (model.id.lower(), provider.id)
if model.enabled:
overrides_by_key[key] = model
else:
disabled_model_keys.add(key)
return overrides_by_key, disabled_model_keys, provider_identities
def _apply_model_visibility(
upstream: BaseUpstreamProvider,
overrides_by_key: dict[ModelKey, ModelRow] | None,
disabled_model_keys: set[ModelKey] | None,
) -> list[object]:
"""Return provider models after DB disabled/override state is applied.
Only the identity fields (``id``, ``forwarded_model_id``,
``canonical_slug``) matter for path discovery, so DB override rows are used
directly rather than rebuilt into fully priced ``Model`` objects — the
pricing pipeline costs ~0.7ms of event-loop CPU per row for data this
module immediately discards.
"""
overrides_by_key = overrides_by_key or {}
disabled_model_keys = disabled_model_keys or set()
upstream_provider_id = getattr(upstream, "db_id", None)
if not isinstance(upstream_provider_id, int):
return [
model
for model in upstream.get_cached_models()
if getattr(model, "enabled", True)
]
visible_models: list[object] = []
seen_model_ids: set[str] = set()
for model in upstream.get_cached_models():
model_id = getattr(model, "id", "")
key = (model_id.lower(), upstream_provider_id)
if not getattr(model, "enabled", True) or key in disabled_model_keys:
continue
# Apply overrides only for this provider's own model row.
override_row = overrides_by_key.get(key)
visible: object = model if override_row is None else override_row
visible_models.append(visible)
seen_model_ids.add(model_id.lower())
# DB-only override rows for this provider with no cached counterpart.
for (model_id_lower, provider_id), override_row in overrides_by_key.items():
if provider_id != upstream_provider_id:
continue
if model_id_lower in seen_model_ids:
continue
visible_models.append(override_row)
seen_model_ids.add(model_id_lower)
return visible_models
async def _collect_provider_paths(
upstream: BaseUpstreamProvider,
provider_identity: ConfiguredProviderIdentity,
overrides_by_key: dict[ModelKey, ModelRow] | None = None,
disabled_model_keys: set[ModelKey] | None = None,
cycle: _RefreshCycleState | None = None,
) -> ProviderPathSnapshot:
"""Collect selectable routes while marking model-level degraded fetches.
A failed OpenRouter lookup preserves only that model's prior rows. Other
models in the same provider still refresh, so a partial outage cannot erase
valid discovery data or freeze the entire provider snapshot.
"""
cycle = cycle or _RefreshCycleState()
models = _apply_model_visibility(upstream, overrides_by_key, disabled_model_keys)
def _base_path(model: object) -> DiscoveredPath:
model_id = exposed_model_id(model)
return DiscoveredPath(
model_id=model_id,
path=encode_model_path(
provider_identity.base_url, provider_identity.id, model_id
),
provider=provider_identity,
)
if not is_openrouter_base_url(upstream.base_url):
return ProviderPathSnapshot(paths=tuple(_base_path(model) for model in models))
if not (upstream.provider_type or "").strip():
return ProviderPathSnapshot(paths=())
semaphore = asyncio.Semaphore(_OPENROUTER_CONCURRENCY)
async with _make_http_client() as client:
async def _for_model(
model: object,
) -> tuple[list[DiscoveredPath], str | None]:
model_id = exposed_model_id(model)
author_slug = openrouter_author_slug(model)
if not author_slug:
return [_base_path(model)], None
endpoints = await _fetch_openrouter_endpoint_subproviders(
client,
upstream.base_url,
upstream.api_key,
author_slug,
semaphore,
cycle,
)
if endpoints is None:
return [], model_id
paths = [_base_path(model)]
paths.extend(
DiscoveredPath(
model_id=model_id,
path=encode_model_path(
provider_identity.base_url,
provider_identity.id,
model_id,
endpoint.tag,
),
provider=provider_identity,
endpoint_tag=endpoint.tag,
endpoint_name=endpoint.provider_name,
)
for endpoint in endpoints
)
return paths, None
results = await asyncio.gather(
*(_for_model(model) for model in models), return_exceptions=True
)
paths: list[DiscoveredPath] = []
preserve_model_ids: set[str] = set()
for model, result in zip(models, results):
if isinstance(result, BaseException):
model_id = exposed_model_id(model)
preserve_model_ids.add(model_id)
logger.warning(
"OpenRouter endpoint discovery task errored",
extra={"provider": upstream.provider_type, "error": str(result)},
)
continue
model_paths, preserved_model_id = result
paths.extend(model_paths)
if preserved_model_id:
preserve_model_ids.add(preserved_model_id)
return ProviderPathSnapshot(
paths=tuple(paths), preserve_model_ids=frozenset(preserve_model_ids)
)
async def _persist_provider_paths(
upstream_provider_id: int, snapshot: ProviderPathSnapshot
) -> None:
"""Replace refreshed rows while retaining model-level degraded snapshots."""
unique_paths = list(
{(path.model_id, path.path): path for path in snapshot.paths}.values()
)
now = int(time.time())
async with create_session() as session:
delete_stmt = delete(ModelPathRow).where(
col(ModelPathRow.upstream_provider_id) == upstream_provider_id
)
if snapshot.preserve_model_ids:
delete_stmt = delete_stmt.where(
col(ModelPathRow.model_id).not_in(sorted(snapshot.preserve_model_ids))
)
await session.exec(delete_stmt) # type: ignore[call-overload]
for start in range(0, len(unique_paths), _PERSIST_CHUNK_SIZE):
chunk = unique_paths[start : start + _PERSIST_CHUNK_SIZE]
values = [
{
"model_id": discovered.model_id,
"path": discovered.path,
"provider_slug": discovered.provider.slug,
"provider_type": discovered.provider.provider_type,
"endpoint_tag": discovered.endpoint_tag,
"endpoint_name": discovered.endpoint_name,
"upstream_provider_id": upstream_provider_id,
"updated_at": now,
}
for discovered in chunk
]
insert_stmt = insert(ModelPathRow).values(values)
await session.execute(
insert_stmt.on_conflict_do_update(
index_elements=["model_id", "path", "upstream_provider_id"],
set_={
"provider_slug": insert_stmt.excluded.provider_slug,
"provider_type": insert_stmt.excluded.provider_type,
"endpoint_tag": insert_stmt.excluded.endpoint_tag,
"endpoint_name": insert_stmt.excluded.endpoint_name,
"updated_at": insert_stmt.excluded.updated_at,
},
)
)
await session.commit()
async def prune_model_paths_for_inactive_providers() -> None:
"""Delete paths whose provider is no longer enabled in the database.
Called from ``refresh_model_maps`` so admin mutations (disable/delete
provider) stop advertising a provider's paths immediately instead of
waiting for the next timed refresh. Uses the DB as the source of truth, so
it is safe at boot even before upstreams initialize.
"""
async with create_session() as session:
enabled_ids = (
await session.exec(
select(UpstreamProviderRow.id).where(
col(UpstreamProviderRow.enabled).is_(True)
)
)
).all()
stmt = delete(ModelPathRow)
if enabled_ids:
stmt = stmt.where(
col(ModelPathRow.upstream_provider_id).not_in(
[pid for pid in enabled_ids if pid is not None]
)
)
await session.exec(stmt) # type: ignore[call-overload]
await session.commit()
async def refresh_model_paths(
upstreams: list[BaseUpstreamProvider],
) -> None:
"""Recompute and persist model paths for every enabled provider.
One provider's failure is logged and isolated; it must not break the rest.
A provider whose paths could not be determined this cycle keeps its
previously persisted rows. An empty ``upstreams`` list (e.g. a failed
``initialize_upstreams`` at boot) is treated as "unknown" and touches
nothing.
"""
if not upstreams:
logger.warning("Skipping model paths refresh: no live upstreams")
return
(
overrides_by_key,
disabled_model_keys,
provider_identities,
) = await _load_model_visibility()
await prune_model_paths_for_inactive_providers()
cycle = _RefreshCycleState()
for upstream in upstreams:
if upstream.db_id is None or upstream.db_id not in provider_identities:
continue
try:
snapshot = await _collect_provider_paths(
upstream,
provider_identity=provider_identities[upstream.db_id],
overrides_by_key=overrides_by_key,
disabled_model_keys=disabled_model_keys,
cycle=cycle,
)
if snapshot.preserve_model_ids:
logger.warning(
"Some model paths are unknown; keeping their previous rows",
extra={
"provider": upstream.provider_type or upstream.base_url,
"db_id": upstream.db_id,
"preserved_models": len(snapshot.preserve_model_ids),
},
)
await _persist_provider_paths(upstream.db_id, snapshot)
except Exception as e: # noqa: BLE001 - isolate per-provider failures
logger.error(
"Failed to refresh model paths for provider",
extra={
"provider": upstream.provider_type or upstream.base_url,
"db_id": upstream.db_id,
"error": str(e),
"error_type": type(e).__name__,
},
)
async def refresh_model_paths_for_provider(upstream_provider_id: int) -> None:
"""Synchronize one provider when model-path discovery is enabled."""
if _refresh_interval_seconds() <= 0:
return
from ..proxy import get_upstreams
matching = [
upstream
for upstream in get_upstreams()
if upstream.db_id == upstream_provider_id
]
if matching:
await refresh_model_paths(matching)
else:
await prune_model_paths_for_inactive_providers()
async def _drain_scheduled_provider_refreshes() -> None:
"""Serialize and coalesce model-path refreshes scheduled by admin writes."""
global _scheduled_provider_refresh_task
try:
# Let mutations in the same event-loop turn collapse into one refresh.
await asyncio.sleep(0)
while _scheduled_provider_refresh_ids:
if _refresh_interval_seconds() <= 0:
_scheduled_provider_refresh_ids.clear()
return
provider_id = min(_scheduled_provider_refresh_ids)
_scheduled_provider_refresh_ids.remove(provider_id)
try:
await refresh_model_paths_for_provider(provider_id)
except asyncio.CancelledError:
raise
except Exception as exc: # noqa: BLE001 - background best effort
logger.warning(
"Failed to refresh model paths after admin mutation",
extra={
"upstream_provider_id": provider_id,
"error": str(exc),
"error_type": type(exc).__name__,
},
)
finally:
_scheduled_provider_refresh_task = None
async def schedule_model_paths_refresh_for_provider(
upstream_provider_id: int,
) -> None:
"""Queue a non-blocking, coalesced refresh after an admin mutation."""
global _scheduled_provider_refresh_task
if _refresh_interval_seconds() <= 0:
return
_scheduled_provider_refresh_ids.add(upstream_provider_id)
if (
_scheduled_provider_refresh_task is None
or _scheduled_provider_refresh_task.done()
):
_scheduled_provider_refresh_task = asyncio.create_task(
_drain_scheduled_provider_refreshes(),
name="model-path-admin-refresh",
)
def _refresh_interval_seconds() -> int:
"""Current interval, re-read every loop so runtime setting changes apply."""
from ..core.settings import settings
if not getattr(settings, "enable_model_paths_refresh", True):
return 0
return int(getattr(settings, "model_paths_refresh_interval_seconds", 0) or 0)
async def refresh_model_paths_periodically(
upstreams_provider: (
Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider]
),
) -> None:
"""Background task mirroring ``refresh_upstreams_models_periodically``.
The interval and enable flag are re-read every iteration, so the refresh
can be turned off (or on) and retuned without a restart. While disabled the
task idles instead of exiting, so re-enabling takes effect.
"""
_DISABLED_POLL_SECONDS = 60.0
def _resolve_upstreams() -> list[BaseUpstreamProvider]:
if callable(upstreams_provider):
return upstreams_provider()
return upstreams_provider
while True:
interval = _refresh_interval_seconds()
if interval <= 0:
try:
await asyncio.sleep(_DISABLED_POLL_SECONDS)
except asyncio.CancelledError:
break
continue
try:
await refresh_model_paths(_resolve_upstreams())
except asyncio.CancelledError:
break
except Exception as e: # noqa: BLE001
logger.error(
"Error in model paths refresh loop",
extra={"error": str(e), "error_type": type(e).__name__},
)
try:
jitter = max(0.0, float(interval) * 0.1)
await asyncio.sleep(interval + random.uniform(0, jitter))
except asyncio.CancelledError:
break
def _serialize_path(row: ModelPathRow) -> dict[str, Any]:
endpoint = None
if row.endpoint_tag or row.endpoint_name:
endpoint = {"tag": row.endpoint_tag, "name": row.endpoint_name}
return {
"path": row.path,
"provider": {
"id": row.upstream_provider_id,
"slug": row.provider_slug,
"type": row.provider_type,
},
"endpoint": endpoint,
}
async def get_all_model_paths() -> dict:
"""All models with their exact selectable routes."""
async with create_session() as session:
rows = (
await session.exec(
select(ModelPathRow).order_by(
col(ModelPathRow.model_id),
col(ModelPathRow.path),
col(ModelPathRow.upstream_provider_id),
)
)
).all()
grouped: dict[str, list[dict[str, Any]]] = {}
seen_paths: dict[str, set[str]] = {}
updated_at = 0
for row in rows:
updated_at = max(updated_at, row.updated_at)
if row.path in seen_paths.setdefault(row.model_id, set()):
continue
seen_paths[row.model_id].add(row.path)
grouped.setdefault(row.model_id, []).append(_serialize_path(row))
data = [
{
"id": grouped_model_id,
"paths": grouped[grouped_model_id],
}
for grouped_model_id in sorted(grouped)
]
return {"data": data, "updated_at": updated_at or None}
async def get_paths_for_model(model_id: str) -> dict:
"""Return paths only for the exact model ID advertised by ``/v1/models``."""
async with create_session() as session:
rows = (
await session.exec(
select(ModelPathRow)
.where(col(ModelPathRow.model_id) == model_id)
.order_by(
col(ModelPathRow.path),
col(ModelPathRow.upstream_provider_id),
)
)
).all()
seen: set[str] = set()
paths: list[dict] = []
updated_at = 0
for row in rows:
updated_at = max(updated_at, row.updated_at)
if row.path in seen:
continue
seen.add(row.path)
paths.append(_serialize_path(row))
return {"data": paths, "updated_at": updated_at or None}
+365 -168
View File
@@ -1,9 +1,14 @@
import asyncio
import fcntl
import os
import re
import socket
import time
import typing
from typing import TypedDict
from contextlib import asynccontextmanager
from contextvars import ContextVar
from pathlib import Path
from typing import AsyncGenerator, TypedDict
import httpx
from cashu.core.base import Proof, Token
@@ -32,6 +37,62 @@ _CashuMintInfo.model_rebuild(force=True)
logger = get_logger(__name__)
# Preserve the real scheduler yield even when payout tests patch asyncio.sleep.
_scheduler_sleep = asyncio.sleep
_WALLET_OPERATION_LOCK = Path(".wallet") / ".routstr-operation.lock"
_wallet_operation_depth: ContextVar[int] = ContextVar(
"wallet_operation_depth", default=0
)
@asynccontextmanager
async def wallet_operation_guard() -> AsyncGenerator[None, None]:
"""Serialize proof mutation and owner payout across local worker processes."""
depth = _wallet_operation_depth.get()
if depth:
token = _wallet_operation_depth.set(depth + 1)
try:
yield
finally:
_wallet_operation_depth.reset(token)
return
_WALLET_OPERATION_LOCK.parent.mkdir(parents=True, exist_ok=True)
fd = os.open(_WALLET_OPERATION_LOCK, os.O_CREAT | os.O_RDWR, 0o600)
acquired = False
depth_token = None
try:
while not acquired:
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
acquired = True
except BlockingIOError:
await _scheduler_sleep(0.05)
depth_token = _wallet_operation_depth.set(1)
yield
finally:
if depth_token is not None:
_wallet_operation_depth.reset(depth_token)
if acquired:
fcntl.flock(fd, fcntl.LOCK_UN)
os.close(fd)
def _sats_to_msats(amount: int) -> int:
return amount * 1000
def _msats_to_sats(amount: int) -> int:
return amount // 1000
def _mints_to_inspect() -> list[str]:
"""Return configured mints plus the primary mint, without duplicates."""
mint_urls = list(settings.cashu_mints)
if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint)
return mint_urls
class MintConnectionError(Exception):
"""The mint could not be reached (network transport failure).
@@ -305,10 +366,10 @@ def _net_minted_amount(amount_msat: int, token_unit: str, fees: int) -> int:
Convert the token value minus fees (given in the token unit) into an
amount in the primary mint's unit.
"""
fee_msat = fees * 1000 if token_unit == "sat" else fees
fee_msat = _sats_to_msats(fees) if token_unit == "sat" else fees
remaining_msat = amount_msat - fee_msat
if settings.primary_mint_unit == "sat":
return int(remaining_msat // 1000)
return _msats_to_sats(remaining_msat)
return int(remaining_msat)
@@ -361,7 +422,7 @@ async def _calculate_swap_amount(
melt fees and NUT-02 input fees on the foreign mint.
"""
if settings.primary_mint_unit == "sat":
receive_amount = amount_msat // 1000
receive_amount = _msats_to_sats(amount_msat)
else:
receive_amount = amount_msat
@@ -395,7 +456,7 @@ async def _calculate_swap_amount(
logger.info(
"swap_to_primary_mint: fee estimation result",
extra={
"token_amount_sat": amount_msat // 1000,
"token_amount_sat": _msats_to_sats(amount_msat),
"estimated_fee": total_fees,
"estimated_fee_unit": token_unit,
"input_fees": input_fees,
@@ -434,7 +495,7 @@ async def swap_to_primary_mint(
token_amount = token_obj.amount
if token_obj.unit == "sat":
amount_msat = token_amount * 1000
amount_msat = _sats_to_msats(token_amount)
elif token_obj.unit == "msat":
amount_msat = token_amount
else:
@@ -598,7 +659,10 @@ async def swap_to_primary_mint(
# advance the counter so the next request derives fresh secrets.
logger.warning(
"swap_to_primary_mint: outputs already signed — recovering orphaned proofs",
extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount},
extra={
"mint_quote_id": mint_quote.quote,
"minted_amount": minted_amount,
},
)
try:
for keyset_id in primary_wallet.keysets:
@@ -667,6 +731,13 @@ async def swap_to_primary_mint(
async def credit_balance(
cashu_token: str, key: db.ApiKey, session: db.AsyncSession
) -> int:
async with wallet_operation_guard():
return await _credit_balance_locked(cashu_token, key, session)
async def _credit_balance_locked(
cashu_token: str, key: db.ApiKey, session: db.AsyncSession
) -> int:
logger.info(
"credit_balance: Starting token redemption",
@@ -683,7 +754,7 @@ async def credit_balance(
)
if unit == "sat":
amount = amount * 1000
amount = _sats_to_msats(amount)
logger.info(
"credit_balance: Converted to msat", extra={"amount_msat": amount}
)
@@ -728,6 +799,9 @@ async def credit_balance(
)
await session.commit()
await session.refresh(key)
# refresh() starts a read transaction; release it before the
# transaction-history write opens its own session below.
await session.commit()
except TokenConsumedError:
raise
except Exception as db_error:
@@ -836,18 +910,34 @@ async def fetch_all_balances(
if units is None:
units = ["sat", "msat"]
async def fetch_balance(
session: db.AsyncSession, mint_url: str, unit: str
) -> BalanceDetail:
try:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
# Received tokens are stored against primary_mint even when cashu_mints is
# empty, so include it in both the liability query and mint fan-out.
mint_urls = _mints_to_inspect()
user_balances: dict[tuple[str, str], int] = {}
liabilities_error: str | None = None
try:
async with db.create_session() as session:
user_balances = await db.balances_by_mint_and_unit(
session, mint_urls, units
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
user_balance = await db.balances_for_mint_and_unit(session, mint_url, unit)
except Exception as e:
logger.error("Error reading user balances", extra={"error": str(e)})
liabilities_error = str(e)
mint_check_limit = asyncio.Semaphore(settings.mint_operation_concurrency)
async def fetch_balance(mint_url: str, unit: str) -> BalanceDetail:
try:
async with mint_check_limit:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
user_balance = user_balances.get((mint_url, unit), 0)
if unit == "sat":
user_balance = user_balance // 1000
user_balance = _msats_to_sats(user_balance)
proofs_balance = sum(proof.amount for proof in proofs)
result: BalanceDetail = {
@@ -870,48 +960,37 @@ async def fetch_all_balances(
}
return error_result
# Build the set of mints to inspect. Received tokens are stored against
# ``primary_mint`` (which defaults to a real mint even when ``cashu_mints``
# is empty), so include it as a fallback — otherwise a node that accepts
# payments would still report empty balances when ``cashu_mints`` is unset.
mint_urls: list[str] = list(settings.cashu_mints)
if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint)
tasks = [fetch_balance(mint_url, unit) for mint_url in mint_urls for unit in units]
balance_details = list(await asyncio.gather(*tasks))
# Create tasks for all mint/unit combinations
async with db.create_session() as session:
tasks = [
fetch_balance(session, mint_url, unit)
for mint_url in mint_urls
for unit in units
]
# Run all tasks concurrently
balance_details = list(await asyncio.gather(*tasks))
# Calculate totals
total_wallet_balance_sats = 0
total_user_balance_sats = 0
for detail in balance_details:
if not detail.get("error"):
# Convert to sats for total calculation
unit = detail["unit"]
proofs_balance_sats = (
detail["wallet_balance"]
if unit == "sat"
else detail["wallet_balance"] // 1000
)
user_balance_sats = (
if detail.get("error"):
continue
unit = detail["unit"]
total_wallet_balance_sats += (
detail["wallet_balance"]
if unit == "sat"
else _msats_to_sats(detail["wallet_balance"])
)
if liabilities_error is None:
total_user_balance_sats += (
detail["user_balance"]
if unit == "sat"
else detail["user_balance"] // 1000
else _msats_to_sats(detail["user_balance"])
)
total_wallet_balance_sats += proofs_balance_sats
total_user_balance_sats += user_balance_sats
owner_balance = total_wallet_balance_sats - total_user_balance_sats
if liabilities_error is None:
owner_balance = total_wallet_balance_sats - total_user_balance_sats
else:
# Custody remains knowable when the DB read fails, but the user/owner
# split does not. Never report unknown liabilities as owner profit.
owner_balance = 0
for detail in balance_details:
detail["user_balance"] = 0
detail["owner_balance"] = 0
detail.setdefault("error", liabilities_error)
return (
balance_details,
@@ -921,6 +1000,70 @@ async def fetch_all_balances(
)
async def _payout_mint_and_unit(mint_url: str, unit: str) -> None:
"""Send only conservatively proven owner funds for one wallet."""
try:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
await asyncio.sleep(5)
except Exception as e:
logger.error(
f"Error sending payout: {type(e).__name__}",
extra={"error": str(e), "mint_url": mint_url, "unit": unit},
)
return
try:
async with db.create_session() as session:
# ApiKey stores a refund preference, not funding provenance. Until
# liabilities have a durable per-credit ledger, subtract the total
# liability from every wallet rather than risk calling customer
# funds owner profit on the wrong mint.
user_balance = await db.total_user_liability(session)
except Exception as e:
logger.error(
f"Error in periodic payout cycle: {type(e).__name__}",
extra={"error": str(e), "mint_url": mint_url, "unit": unit},
)
return
try:
if unit == "sat":
user_balance = _msats_to_sats(user_balance)
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
min_amount = (
settings.min_payout_sat
if unit == "sat"
else _sats_to_msats(settings.min_payout_sat)
)
if available_balance > min_amount:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
settings.receive_ln_address,
unit,
amount=available_balance,
)
logger.info(
"Payout sent successfully",
extra={
"mint_url": mint_url,
"unit": unit,
"balance": available_balance,
"amount_received": amount_received,
},
)
except Exception as e:
logger.error(
f"Error sending payout: {type(e).__name__}",
extra={"error": str(e), "mint_url": mint_url, "unit": unit},
)
async def periodic_payout() -> None:
while True:
await asyncio.sleep(settings.payout_interval_seconds)
@@ -928,64 +1071,13 @@ async def periodic_payout() -> None:
if not settings.receive_ln_address:
continue
# Include the primary mint even if it is not listed in cashu_mints,
# matching fetch_all_balances(); otherwise primary-mint funds never
# auto-payout.
mint_urls: list[str] = list(settings.cashu_mints)
if settings.primary_mint and settings.primary_mint not in mint_urls:
mint_urls.append(settings.primary_mint)
async with db.create_session() as session:
for mint_url in mint_urls:
for unit in ["sat", "msat"]:
# Isolate failures per mint/unit so one slow or failing
# mint does not abort payout for every other mint/unit.
try:
wallet = await get_wallet(mint_url, unit)
proofs = get_proofs_per_mint_and_unit(
wallet, mint_url, unit, not_reserved=True
)
proofs = await slow_filter_spend_proofs(proofs, wallet)
await asyncio.sleep(5)
user_balance = await db.balances_for_mint_and_unit(
session, mint_url, unit
)
if unit == "sat":
user_balance = user_balance // 1000
proofs_balance = sum(proof.amount for proof in proofs)
available_balance = proofs_balance - user_balance
# Threshold is configured in sats; convert for msat wallets.
min_amount = (
settings.min_payout_sat
if unit == "sat"
else settings.min_payout_sat * 1000
)
if available_balance > min_amount:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
settings.receive_ln_address,
unit,
amount=available_balance,
)
logger.info(
"Payout sent successfully",
extra={
"mint_url": mint_url,
"unit": unit,
"balance": available_balance,
"amount_received": amount_received,
},
)
except Exception as e:
logger.error(
f"Error sending payout: {type(e).__name__}",
extra={
"error": str(e),
"mint_url": mint_url,
"unit": unit,
},
)
for mint_url in _mints_to_inspect():
for unit in ["sat", "msat"]:
# Proof mutation, liability observation, and sending are one
# cross-process critical section. A credit takes the same
# lock from before redemption through its liability commit.
async with wallet_operation_guard():
await _payout_mint_and_unit(mint_url, unit)
except Exception as e:
logger.error(
f"Error in periodic payout cycle: {type(e).__name__}",
@@ -993,22 +1085,65 @@ async def periodic_payout() -> None:
)
async def _set_refund_sweep_state(
refund_id: str,
*,
predicates: tuple[typing.Any, ...] = (),
**values: object,
) -> int:
async with db.create_session() as session:
result = await session.exec( # type: ignore[call-overload]
update(db.CashuTransaction)
.where(col(db.CashuTransaction.id) == refund_id, *predicates)
.values(**values)
)
await session.commit()
return int(result.rowcount or 0)
async def _refund_sweep_once(cutoff: int) -> None:
claim_cutoff = int(time.time()) - settings.refund_sweep_claim_timeout_seconds
claim_available = col(db.CashuTransaction.sweep_started_at).is_(None) | (
col(db.CashuTransaction.sweep_started_at) < claim_cutoff
)
async with db.create_session() as session:
stmt = select(db.CashuTransaction).where(
db.CashuTransaction.type == "out",
db.CashuTransaction.collected == False, # noqa: E712
db.CashuTransaction.swept == False, # noqa: E712
db.CashuTransaction.created_at < cutoff,
claim_available,
)
results = await session.exec(stmt)
refunds = results.all()
for refund in refunds:
try:
await recieve_token(refund.token)
refund.swept = True
session.add(refund)
for refund in refunds:
reclaimed_stale_claim = refund.sweep_started_at is not None
claim_started_at = int(time.time())
claimed = await _set_refund_sweep_state(
refund.id,
predicates=(
col(db.CashuTransaction.swept) == False, # noqa: E712
col(db.CashuTransaction.collected) == False, # noqa: E712
claim_available,
),
sweep_started_at=claim_started_at,
)
if claimed != 1:
continue
claim_owned = col(db.CashuTransaction.sweep_started_at) == claim_started_at
redeemed = False
try:
await recieve_token(refund.token)
redeemed = True
finalized = await _set_refund_sweep_state(
refund.id,
predicates=(claim_owned,),
swept=True,
sweep_started_at=None,
)
if finalized == 1:
logger.info(
"Swept uncollected refund",
extra={
@@ -1017,26 +1152,70 @@ async def _refund_sweep_once(cutoff: int) -> None:
"unit": refund.unit,
},
)
except Exception as e:
error_msg = str(e).lower()
if "already spent" in error_msg:
refund.collected = True
session.add(refund)
else:
logger.critical(
"Refund token swept after claim ownership changed; manual reconciliation required",
extra={"id": refund.id},
)
except BaseException as e:
if redeemed or isinstance(e, TokenConsumedError):
# The token was spent, or the redemption outcome is known to be
# post-spend. Retain the claim so a stale retry classifies
# "already spent" as swept, never as a client collection.
logger.critical(
"Refund token spent but sweep checkpoint was not completed; manual reconciliation required",
extra={"id": refund.id},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
continue
error_msg = str(e).lower()
if isinstance(e, Exception) and "already spent" in error_msg:
if reclaimed_stale_claim:
# A prior worker may have redeemed the token and crashed
# before finalizing. Treat the ambiguous stale claim as a
# completed sweep rather than misreporting client collection.
updated = await _set_refund_sweep_state(
refund.id,
predicates=(claim_owned,),
swept=True,
sweep_started_at=None,
)
else:
updated = await _set_refund_sweep_state(
refund.id,
predicates=(claim_owned,),
collected=True,
swept=False,
sweep_started_at=None,
)
if updated == 1:
logger.info(
"Refund already spent (client collected), marking swept",
"Refund token was already spent",
extra={
"id": refund.id,
"reclaimed_stale_claim": reclaimed_stale_claim,
},
)
else:
logger.warning(
"Failed to sweep refund",
extra={
"id": refund.id,
"error": str(e),
},
"Refund claim ownership changed before spent-token checkpoint",
extra={"id": refund.id},
)
await session.commit()
else:
# Once redemption starts, an exception cannot prove the token
# was not spent (for example, a melt may land before the
# response is lost). Retain the claim so a stale retry treats
# an "already spent" result as a completed sweep.
logger.critical(
"Refund token redemption outcome is unknown; retaining sweep claim for reconciliation",
extra={"id": refund.id, "error": str(e)},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
async def refund_sweep_once() -> None:
@@ -1082,53 +1261,71 @@ async def periodic_routstr_fee_payout() -> None:
)
continue
accumulated_sats = fee.accumulated_msats // 1000
if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT:
wallet = await get_wallet(settings.primary_mint, "sat")
proofs = get_proofs_per_mint_and_unit(
wallet, settings.primary_mint, "sat", not_reserved=True
)
paid_msats = accumulated_sats * 1000
payout_checkpointed = await db.reset_routstr_fee(
session, paid_msats
)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
accumulated_sats = _msats_to_sats(fee.accumulated_msats)
if accumulated_sats < ROUTSTR_FEE_DEFAULT_PAYOUT:
continue
paid_msats = _sats_to_msats(accumulated_sats)
try:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
ROUTSTR_LN_ADDRESS,
"sat",
amount=accumulated_sats,
)
except Exception:
logger.critical(
"Routstr fee payout outcome is unknown; manual reconciliation required",
extra={"payout_in_progress_msats": paid_msats},
exc_info=True,
)
continue
# Wallet/proof preparation cannot send funds, so do it before the
# durable checkpoint. A preparation failure must not strand an
# in-progress payout that requires manual reconciliation.
wallet = await get_wallet(settings.primary_mint, "sat")
proofs = get_proofs_per_mint_and_unit(
wallet, settings.primary_mint, "sat", not_reserved=True
)
async with db.create_session() as session:
payout_checkpointed = await db.reset_routstr_fee(session, paid_msats)
if not payout_checkpointed:
logger.warning("Routstr fee payout was already claimed")
continue
try:
amount_received = await raw_send_to_lnurl(
wallet,
proofs,
ROUTSTR_LN_ADDRESS,
"sat",
amount=accumulated_sats,
)
except BaseException as e:
logger.critical(
"Routstr fee payout outcome is unknown; manual reconciliation required",
extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
continue
try:
async with db.create_session() as session:
payout_completed = await db.complete_routstr_fee_payout(
session, paid_msats
)
if not payout_completed:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
)
continue
except BaseException as e:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
exc_info=isinstance(e, Exception),
)
if not isinstance(e, Exception):
raise
continue
if not payout_completed:
logger.critical(
"Routstr fee payout sent but checkpoint was not completed",
extra={"payout_in_progress_msats": paid_msats},
)
continue
logger.info(
"Routstr fee payout sent",
extra={
"accumulated_sats": accumulated_sats,
"amount_received": amount_received,
},
)
logger.info(
"Routstr fee payout sent",
extra={
"accumulated_sats": accumulated_sats,
"amount_received": amount_received,
},
)
except Exception as e:
logger.error(
f"Error in Routstr fee payout: {type(e).__name__}",
@@ -3,20 +3,23 @@
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 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 _make_invoice(**kwargs: object) -> LightningInvoice:
@@ -48,6 +51,7 @@ def mock_wallet_mint() -> object:
# Persistence
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_invoice_persists_balance_limit(
integration_session: AsyncSession,
@@ -92,6 +96,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 +105,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 +121,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 +138,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 +146,243 @@ 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()
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")
return []
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()
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()
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
wallet.mint = AsyncMock(return_value=[])
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 +391,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 +399,84 @@ 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()
wallet.get_mint_quote = AsyncMock(return_value=MagicMock(paid=True))
async def always_succeeding_mint(*args: object, **kwargs: object) -> list[object]:
await asyncio.sleep(0.05)
return []
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 == 2
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
@@ -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
+35
View File
@@ -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]
+72
View File
@@ -221,6 +221,78 @@ def _make_api_key(
return key
@pytest.mark.asyncio
async def test_apikey_refund_returns_persisted_token_after_cache_loss() -> None:
key = _make_api_key(balance=0, refund_currency="sat")
refund_token = "cashuApersisted_refund_token"
refund_tx = _make_cashu_tx(
token=refund_token,
amount=5,
unit="sat",
type="out",
request_id=None,
)
refund_tx.source = "apikey"
refund_tx.api_key_hashed_key = key.hashed_key
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=_exec_result(refund_tx))
session.add = MagicMock()
session.commit = AsyncMock()
with (
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
patch("routstr.balance.send_token", AsyncMock()) as mock_send_token,
):
result = await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert result == {"token": refund_token, "sats": "5"}
assert refund_tx.collected is True
session.add.assert_called_once_with(refund_tx)
session.commit.assert_awaited_once()
mock_send_token.assert_not_awaited()
@pytest.mark.asyncio
async def test_apikey_refund_rejects_persisted_token_after_sweep() -> None:
from fastapi import HTTPException
key = _make_api_key(balance=0, refund_currency="sat")
refund_tx = _make_cashu_tx(
token="cashuAswept_apikey_refund",
amount=5,
unit="sat",
request_id=None,
swept=True,
)
refund_tx.source = "apikey"
refund_tx.api_key_hashed_key = key.hashed_key
session = MagicMock()
session.get = AsyncMock(return_value=key)
session.exec = AsyncMock(return_value=_exec_result(refund_tx))
session.add = MagicMock()
session.commit = AsyncMock()
with patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)):
with pytest.raises(HTTPException) as exc_info:
await refund_wallet_endpoint(
authorization="Bearer sk-testhash",
x_cashu=None,
session=session,
)
assert exc_info.value.status_code == 410
assert exc_info.value.detail == "Refund has been swept"
session.add.assert_not_called()
session.commit.assert_not_awaited()
@pytest.mark.asyncio
async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> None:
key = _make_api_key(balance=5000, refund_currency="sat")
@@ -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"], []) == {}
+85
View File
@@ -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()
+238 -10
View File
@@ -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()
+1 -7
View File
@@ -4,9 +4,6 @@ import subprocess
import sys
from pathlib import Path
from alembic.config import Config
from alembic.script import ScriptDirectory
def _run_alembic(root: Path, database_url: str, revision: str) -> None:
env = os.environ.copy()
@@ -40,10 +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()
migration_config = Config(str(root / "alembic.ini"))
assert version == (
ScriptDirectory.from_config(migration_config).get_current_head(),
)
assert version == ("aa50fde387a2",)
assert {
"id",
"accumulated_msats",
+220 -9
View File
@@ -1,7 +1,13 @@
import asyncio
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch
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
@@ -26,8 +32,10 @@ def _patches( # type: ignore[no-untyped-def]
AsyncMock(side_effect=lambda proofs, wallet: 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 +46,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 +63,159 @@ 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_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 +239,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 +254,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"
File diff suppressed because it is too large Load Diff
+123 -47
View File
@@ -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"]
+281 -5
View File
@@ -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
+95
View File
@@ -62,6 +62,101 @@ def test_payout_settings_have_sensible_defaults() -> None:
assert s.payout_interval_seconds == 900
def test_database_pool_defaults_provide_concurrency_headroom() -> None:
s = Settings()
assert s.database_pool_size == 10
assert s.database_max_overflow == 20
assert s.database_pool_timeout == 15.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", 10)
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 == 10
# ...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",
[