mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-31 07:46:15 +00:00
Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2268c3a962 | ||
|
|
c4d27ba02a | ||
|
|
5ea5024608 | ||
|
|
39f801561b | ||
|
|
ff55788e2d | ||
|
|
1131c2d583 | ||
|
|
7108d554c8 | ||
|
|
1eddf89d52 | ||
|
|
a2cedd6769 | ||
|
|
7b0ade3987 | ||
|
|
2410a4a6ce | ||
|
|
1b09639265 | ||
|
|
7394b10e75 |
@@ -22,6 +22,19 @@ ROUTSTR_SECRET_KEY=
|
||||
|
||||
# Database
|
||||
# DATABASE_URL=sqlite+aiosqlite:///keys.db
|
||||
# Pool controls are validated at boot, sourced only from the environment, and
|
||||
# logged at startup. Keep total capacity across all workers below the database
|
||||
# connection limit. Pre-ping is automatic for networked backends; SQLite may
|
||||
# explicitly opt in if desired.
|
||||
# DATABASE_POOL_SIZE=5
|
||||
# DATABASE_MAX_OVERFLOW=10
|
||||
# DATABASE_POOL_TIMEOUT=30
|
||||
# DATABASE_POOL_RECYCLE=1800
|
||||
# DATABASE_POOL_PRE_PING=false
|
||||
# Warn when a checkout is held this many seconds.
|
||||
# DATABASE_POOL_HOLD_WARN_SECONDS=10
|
||||
# SQLite serialises writes; increasing its pool can trade pool timeouts for
|
||||
# "database is locked" errors rather than increasing write throughput.
|
||||
|
||||
# Node Information
|
||||
# NAME=My Routstr Node
|
||||
@@ -31,7 +44,9 @@ ROUTSTR_SECRET_KEY=
|
||||
# RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com"
|
||||
# ENABLE_ANALYTICS_SHARING=true
|
||||
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me"
|
||||
# MINT_OPERATION_CONCURRENCY=4
|
||||
# RECEIVE_LN_ADDRESS=
|
||||
# REFUND_SWEEP_CLAIM_TIMEOUT_SECONDS=900
|
||||
|
||||
# Custom Pricing Configuration
|
||||
# MODEL_BASED_PRICING=true
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
@@ -0,0 +1,35 @@
|
||||
"""repair refund sweep claim lease
|
||||
|
||||
Revision ID: b7c9d1e3f5a2
|
||||
Revises: aa50fde387a2
|
||||
Create Date: 2026-07-30 03:10:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "b7c9d1e3f5a2"
|
||||
down_revision = "aa50fde387a2"
|
||||
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:
|
||||
# Revision aa50fde387a2 owns this column, so downgrading only the repair
|
||||
# migration must preserve the schema expected at that revision.
|
||||
pass
|
||||
+8
-16
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
+109
-64
@@ -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,43 @@ async def complete_routstr_fee_payout(session: AsyncSession, paid_msats: int) ->
|
||||
return result.rowcount == 1
|
||||
|
||||
|
||||
async def balances_for_mint_and_unit(
|
||||
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
@@ -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
@@ -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 match
|
||||
# SQLAlchemy's established queue-pool behavior. Pre-ping is enabled by the
|
||||
# engine factory for networked backends; SQLite can explicitly opt in.
|
||||
# These fields are env-only below.
|
||||
database_pool_size: int = Field(default=5, ge=1, env="DATABASE_POOL_SIZE")
|
||||
database_max_overflow: int = Field(default=10, ge=0, env="DATABASE_MAX_OVERFLOW")
|
||||
database_pool_timeout: float = Field(
|
||||
default=30.0, gt=0, env="DATABASE_POOL_TIMEOUT"
|
||||
)
|
||||
database_pool_recycle: int = Field(default=1800, ge=0, env="DATABASE_POOL_RECYCLE")
|
||||
database_pool_pre_ping: bool = Field(default=False, env="DATABASE_POOL_PRE_PING")
|
||||
database_pool_hold_warn_seconds: float = Field(
|
||||
default=10.0, gt=0, env="DATABASE_POOL_HOLD_WARN_SECONDS"
|
||||
)
|
||||
|
||||
# Logging
|
||||
@@ -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
|
||||
|
||||
+132
-41
@@ -5,7 +5,8 @@ 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
|
||||
@@ -222,80 +223,170 @@ async def recover_invoice(
|
||||
async def check_invoice_payment(
|
||||
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
|
||||
|
||||
@@ -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")
|
||||
|
||||
+15
-13
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
@@ -183,19 +184,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 +221,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
|
||||
|
||||
@@ -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}
|
||||
+310
-155
@@ -33,6 +33,22 @@ _CashuMintInfo.model_rebuild(force=True)
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
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 +321,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 +377,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 +411,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 +450,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 +614,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:
|
||||
@@ -683,7 +702,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}
|
||||
)
|
||||
@@ -836,18 +855,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 +905,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,
|
||||
@@ -931,61 +955,87 @@ async def periodic_payout() -> None:
|
||||
# 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)
|
||||
mint_urls = _mints_to_inspect()
|
||||
|
||||
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(
|
||||
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)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending payout: {type(e).__name__}",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"mint_url": mint_url,
|
||||
"unit": unit,
|
||||
},
|
||||
)
|
||||
continue
|
||||
|
||||
# Fetch the liability AFTER the proofs snapshot and the
|
||||
# settle delay: a concurrent top-up then only inflates the
|
||||
# liability, shrinking the payout — never sending
|
||||
# customer-backed funds as profit.
|
||||
try:
|
||||
async with db.create_session() as session:
|
||||
user_balance = await db.balance_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
|
||||
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,
|
||||
},
|
||||
)
|
||||
continue
|
||||
|
||||
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
|
||||
# Threshold is configured in sats; convert for msat wallets.
|
||||
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,
|
||||
)
|
||||
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__}",
|
||||
logger.info(
|
||||
"Payout sent successfully",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"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,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error in periodic payout cycle: {type(e).__name__}",
|
||||
@@ -993,22 +1043,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 +1110,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 +1219,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,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]
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Real-DB coverage for db.balances_by_mint_and_unit.
|
||||
|
||||
Verifies the grouped liability query used by fetch_all_balances: it sums
|
||||
balances per (mint_url, unit), filters to the requested mints/units, excludes
|
||||
NULL mint/currency rows, and returns nothing for empty inputs.
|
||||
"""
|
||||
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlalchemy.pool import StaticPool
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import (
|
||||
ApiKey,
|
||||
balance_for_mint_and_unit,
|
||||
balances_by_mint_and_unit,
|
||||
)
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
return create_async_engine(
|
||||
"sqlite+aiosqlite://",
|
||||
poolclass=StaticPool,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def session() -> "AsyncGenerator[AsyncSession, None]":
|
||||
engine = _make_engine()
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
db_session = AsyncSession(engine, expire_on_commit=False)
|
||||
try:
|
||||
yield db_session
|
||||
finally:
|
||||
await db_session.close()
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _add_key(
|
||||
session: AsyncSession,
|
||||
hashed_key: str,
|
||||
balance: int,
|
||||
mint_url: str | None,
|
||||
currency: str | None,
|
||||
) -> None:
|
||||
session.add(
|
||||
ApiKey(
|
||||
hashed_key=hashed_key,
|
||||
balance=balance,
|
||||
refund_mint_url=mint_url,
|
||||
refund_currency=currency,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sums_and_groups_by_mint_and_unit(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 500, "http://m1", "sat")
|
||||
await _add_key(session, "c", 7000, "http://m1", "msat")
|
||||
await _add_key(session, "d", 200, "http://m2", "sat")
|
||||
|
||||
result = await balances_by_mint_and_unit(
|
||||
session, ["http://m1", "http://m2"], ["sat", "msat"]
|
||||
)
|
||||
|
||||
assert result[("http://m1", "sat")] == 1500
|
||||
assert result[("http://m1", "msat")] == 7000
|
||||
assert result[("http://m2", "sat")] == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filters_out_unrequested_mints_and_units(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://wanted", "sat")
|
||||
await _add_key(session, "b", 999, "http://other", "sat")
|
||||
await _add_key(session, "c", 888, "http://wanted", "usd")
|
||||
|
||||
result = await balances_by_mint_and_unit(session, ["http://wanted"], ["sat"])
|
||||
|
||||
assert result == {("http://wanted", "sat"): 1000}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_excludes_rows_with_null_mint_or_currency(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 4242, None, None)
|
||||
|
||||
result = await balances_by_mint_and_unit(session, ["http://m1"], ["sat"])
|
||||
|
||||
assert result == {("http://m1", "sat"): 1000}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scalar_balance_for_one_mint_and_unit(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
await _add_key(session, "b", 500, "http://m1", "sat")
|
||||
await _add_key(session, "c", 9000, "http://m1", "msat")
|
||||
await _add_key(session, "d", 700, "http://m2", "sat")
|
||||
|
||||
assert await balance_for_mint_and_unit(session, "http://m1", "sat") == 1500
|
||||
assert await balance_for_mint_and_unit(session, "http://missing", "sat") == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_inputs_return_empty_mapping(session: AsyncSession) -> None:
|
||||
await _add_key(session, "a", 1000, "http://m1", "sat")
|
||||
|
||||
assert await balances_by_mint_and_unit(session, [], ["sat"]) == {}
|
||||
assert await balances_by_mint_and_unit(session, ["http://m1"], []) == {}
|
||||
@@ -0,0 +1,85 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
from routstr.core import db
|
||||
from routstr.core.db import create_db_engine
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_uses_validated_bounded_pool_settings(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: object
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_size", 12)
|
||||
monkeypatch.setattr(settings, "database_max_overflow", 3)
|
||||
monkeypatch.setattr(settings, "database_pool_timeout", 2.5)
|
||||
monkeypatch.setattr(settings, "database_pool_recycle", 900)
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
|
||||
engine = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/pool.db")
|
||||
try:
|
||||
assert engine.pool.size() == 12 # type: ignore[attr-defined]
|
||||
assert engine.pool._max_overflow == 3 # type: ignore[attr-defined]
|
||||
assert engine.pool._timeout == 2.5 # type: ignore[attr-defined]
|
||||
assert engine.pool._recycle == 900
|
||||
assert engine.pool._pre_ping is False
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_memory_sqlite_keeps_static_pool(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", True)
|
||||
engine = create_db_engine("sqlite+aiosqlite://")
|
||||
try:
|
||||
assert isinstance(engine.pool, StaticPool)
|
||||
assert engine.pool._pre_ping is True
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def test_non_sqlite_backend_enables_pre_ping_automatically(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
fake_engine = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(db, "create_async_engine", return_value=fake_engine) as factory,
|
||||
patch.object(db.event, "listen") as listen,
|
||||
):
|
||||
created = create_db_engine("postgresql+asyncpg://user:pass@db/node")
|
||||
|
||||
assert created is fake_engine
|
||||
assert factory.call_args.kwargs["pool_pre_ping"] is True
|
||||
assert listen.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_created_engine_warns_for_long_checkouts(
|
||||
monkeypatch: pytest.MonkeyPatch, tmp_path: object
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "database_pool_hold_warn_seconds", 0.0)
|
||||
monkeypatch.setattr(settings, "database_pool_pre_ping", False)
|
||||
first = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/first.db")
|
||||
second = create_db_engine(f"sqlite+aiosqlite:///{tmp_path}/second.db")
|
||||
|
||||
try:
|
||||
with patch.object(db.logger, "warning") as warning:
|
||||
async with first.connect() as connection:
|
||||
await connection.exec_driver_sql("SELECT 1")
|
||||
async with second.connect() as connection:
|
||||
await connection.exec_driver_sql("SELECT 1")
|
||||
|
||||
assert warning.call_count == 2
|
||||
assert all(
|
||||
call.kwargs["extra"]["threshold_seconds"] == 0.0
|
||||
for call in warning.call_args_list
|
||||
)
|
||||
finally:
|
||||
await first.dispose()
|
||||
await second.dispose()
|
||||
@@ -1,5 +1,5 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
@@ -13,9 +13,19 @@ from routstr import wallet
|
||||
from routstr.core import db
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _session_context(session: Mock) -> AsyncIterator[Mock]:
|
||||
yield session
|
||||
class _SessionContext:
|
||||
def __init__(self, session: Mock) -> None:
|
||||
self.session = session
|
||||
|
||||
async def __aenter__(self) -> Mock:
|
||||
return self.session
|
||||
|
||||
async def __aexit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _session_context(session: Mock) -> _SessionContext:
|
||||
return _SessionContext(session)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -47,7 +57,7 @@ async def test_fee_payout_checkpoint_is_atomic_and_durable() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
async def test_fee_payout_prepares_wallet_then_checkpoints_before_sending() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
@@ -57,6 +67,10 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
payout_wallet = Mock()
|
||||
events: list[str] = []
|
||||
|
||||
async def prepare(*_args: object) -> Mock:
|
||||
events.append("prepare")
|
||||
return payout_wallet
|
||||
|
||||
async def checkpoint(*_args: object) -> bool:
|
||||
events.append("checkpoint")
|
||||
return True
|
||||
@@ -77,18 +91,92 @@ async def test_fee_payout_checkpoints_before_sending() -> None:
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", side_effect=checkpoint),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", side_effect=complete),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=payout_wallet)),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(side_effect=prepare)),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
assert events == ["checkpoint", "send", "complete"]
|
||||
assert events == ["prepare", "checkpoint", "send", "complete"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_preparation_failure_does_not_checkpoint() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
checkpoint = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", checkpoint),
|
||||
patch(
|
||||
"routstr.wallet.get_wallet",
|
||||
AsyncMock(side_effect=RuntimeError("wallet unavailable")),
|
||||
),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
checkpoint.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_lost_checkpoint_race_does_not_send() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
send = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch(
|
||||
"routstr.wallet.db.reset_routstr_fee",
|
||||
AsyncMock(return_value=False),
|
||||
),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", send),
|
||||
patch("routstr.wallet.logger.warning") as warning,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
send.assert_not_awaited()
|
||||
warning.assert_called_once_with("Routstr fee payout was already claimed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -107,7 +195,9 @@ async def test_fee_payout_does_not_retry_an_unresolved_checkpoint() -> None:
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock()) as checkpoint,
|
||||
patch("routstr.wallet.get_wallet", AsyncMock()) as get_wallet,
|
||||
@@ -141,7 +231,9 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", return_value=_session_context(session)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
|
||||
@@ -158,3 +250,139 @@ async def test_fee_payout_keeps_checkpoint_when_send_outcome_is_unknown() -> Non
|
||||
|
||||
complete.assert_not_awaited()
|
||||
critical.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_cancellation_during_send_alerts_and_propagates() -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
complete = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch("routstr.wallet.asyncio.sleep", AsyncMock(return_value=None)),
|
||||
patch(
|
||||
"routstr.wallet.db.create_session", return_value=_session_context(session)
|
||||
),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", complete),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch(
|
||||
"routstr.wallet.raw_send_to_lnurl",
|
||||
AsyncMock(side_effect=asyncio.CancelledError()),
|
||||
),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
complete.assert_not_awaited()
|
||||
critical.assert_called_once()
|
||||
assert critical.call_args.args[0] == (
|
||||
"Routstr fee payout outcome is unknown; manual reconciliation required"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure_site", ["session", "completion"])
|
||||
async def test_fee_payout_completion_failures_use_sent_checkpoint_alert(
|
||||
failure_site: str,
|
||||
) -> None:
|
||||
session = Mock()
|
||||
fee = SimpleNamespace(
|
||||
accumulated_msats=5_000,
|
||||
payout_in_progress_msats=0,
|
||||
payout_started_at=None,
|
||||
)
|
||||
completion = AsyncMock()
|
||||
if failure_site == "session":
|
||||
create_session = Mock(
|
||||
side_effect=[
|
||||
_session_context(session),
|
||||
_session_context(session),
|
||||
RuntimeError("pool unavailable"),
|
||||
]
|
||||
)
|
||||
else:
|
||||
create_session = Mock(return_value=_session_context(session))
|
||||
completion.side_effect = RuntimeError("checkpoint unavailable")
|
||||
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", create_session),
|
||||
patch("routstr.wallet.db.get_routstr_fee", AsyncMock(return_value=fee)),
|
||||
patch("routstr.wallet.db.reset_routstr_fee", AsyncMock(return_value=True)),
|
||||
patch("routstr.wallet.db.complete_routstr_fee_payout", completion),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", AsyncMock(return_value=5)),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
critical.assert_called_once()
|
||||
assert critical.call_args.args[0] == (
|
||||
"Routstr fee payout sent but checkpoint was not completed"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_releases_db_connection_during_send(tmp_path: object) -> None:
|
||||
"""With pool_size=1, the payout must not hold a connection while the
|
||||
external LNURL send is in flight, or the completion step would starve."""
|
||||
engine = create_async_engine(
|
||||
f"sqlite+aiosqlite:///{tmp_path}/payout.db", pool_size=1, max_overflow=0
|
||||
)
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(SQLModel.metadata.create_all)
|
||||
async with AsyncSession(engine) as session:
|
||||
session.add(db.RoutstrFee(id=1, accumulated_msats=5_000_000))
|
||||
await session.commit()
|
||||
|
||||
@asynccontextmanager
|
||||
async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
|
||||
async def send(*_args: object, **_kwargs: object) -> int:
|
||||
assert engine.pool.checkedout() == 0 # type: ignore[attr-defined]
|
||||
return 5
|
||||
|
||||
try:
|
||||
with (
|
||||
patch("routstr.auth.ROUTSTR_FEE_DEFAULT_PAYOUT", 1),
|
||||
patch("routstr.auth.ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS", 1),
|
||||
patch("routstr.auth.ROUTSTR_LN_ADDRESS", "fees@example.com"),
|
||||
patch(
|
||||
"routstr.wallet.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
),
|
||||
patch("routstr.wallet.db.create_session", create_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=Mock())),
|
||||
patch("routstr.wallet.get_proofs_per_mint_and_unit", return_value=[]),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", side_effect=send),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wallet.periodic_routstr_fee_payout()
|
||||
|
||||
async with AsyncSession(engine) as session:
|
||||
fee = await db.get_routstr_fee(session)
|
||||
assert fee.payout_in_progress_msats == 0
|
||||
assert fee.total_paid_msats == 5_000_000
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
@@ -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 == ("b7c9d1e3f5a2",)
|
||||
assert {
|
||||
"id",
|
||||
"accumulated_msats",
|
||||
|
||||
@@ -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
@@ -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.balance_for_mint_and_unit",
|
||||
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.balance_for_mint_and_unit",
|
||||
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.balance_for_mint_and_unit",
|
||||
AsyncMock(return_value=0),
|
||||
),
|
||||
patch("routstr.wallet.raw_send_to_lnurl", raw_send),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
@@ -128,27 +192,39 @@ async def test_periodic_payout_isolates_failing_mint() -> None:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_payout_handles_session_creation_failure() -> None:
|
||||
"""A db.create_session failure is logged and the payout loop continues."""
|
||||
"""A db.create_session failure is logged per mint/unit and the loop continues."""
|
||||
from routstr.core.settings import settings
|
||||
|
||||
create_session = MagicMock(side_effect=RuntimeError("db unavailable"))
|
||||
logger = MagicMock()
|
||||
|
||||
with patch.object(settings, "cashu_mints", ["http://mint:3338"]), patch.object(
|
||||
settings, "primary_mint", "http://mint:3338"
|
||||
), patch.object(settings, "receive_ln_address", "owner@ln.tld"), patch.object(
|
||||
settings, "payout_interval_seconds", _INTERVAL
|
||||
), patch(
|
||||
"routstr.wallet.asyncio.sleep", _one_cycle_sleep()
|
||||
), patch(
|
||||
"routstr.wallet.db.create_session", create_session
|
||||
), patch("routstr.wallet.logger", logger):
|
||||
with (
|
||||
patch.object(settings, "cashu_mints", ["http://mint:3338"]),
|
||||
patch.object(settings, "primary_mint", "http://mint:3338"),
|
||||
patch.object(settings, "receive_ln_address", "owner@ln.tld"),
|
||||
patch.object(settings, "payout_interval_seconds", _INTERVAL),
|
||||
patch("routstr.wallet.asyncio.sleep", _one_cycle_sleep()),
|
||||
patch("routstr.wallet.db.create_session", create_session),
|
||||
patch("routstr.wallet.get_wallet", AsyncMock(return_value=MagicMock())),
|
||||
patch(
|
||||
"routstr.wallet.get_proofs_per_mint_and_unit",
|
||||
MagicMock(return_value=[MagicMock(amount=100_000)]),
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet.slow_filter_spend_proofs",
|
||||
AsyncMock(side_effect=lambda proofs, wallet: proofs),
|
||||
),
|
||||
patch("routstr.wallet.logger", logger),
|
||||
):
|
||||
with pytest.raises(_LoopBreak):
|
||||
await periodic_payout()
|
||||
|
||||
create_session.assert_called_once()
|
||||
logger.error.assert_called_once()
|
||||
# The liability session is opened per mint/unit (sat + msat), and each
|
||||
# DB failure retains the cycle-specific alert wording while remaining
|
||||
# isolated to its own iteration.
|
||||
assert create_session.call_count == 2
|
||||
assert logger.error.call_count == 2
|
||||
message = logger.error.call_args.args[0]
|
||||
extra = logger.error.call_args.kwargs["extra"]
|
||||
assert message == "Error in periodic payout cycle: RuntimeError"
|
||||
assert extra == {"error": "db unavailable"}
|
||||
assert extra["error"] == "db unavailable"
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
from collections.abc import AsyncIterator
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_closes_request_session_before_returning_response() -> None:
|
||||
"""Route completion must release DB resources before response delivery."""
|
||||
request = MagicMock()
|
||||
request.method = "GET"
|
||||
request.headers = {"accept": "application/json"}
|
||||
request.url.path = "/not-an-api-route"
|
||||
request.state.request_id = "test-request"
|
||||
session = AsyncMock()
|
||||
|
||||
response = await proxy_module.proxy(request, "not-an-api-route", session=session)
|
||||
|
||||
assert response.status_code == 404
|
||||
session.close.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_session_is_closed_before_first_stream_chunk() -> None:
|
||||
request = MagicMock()
|
||||
session = AsyncMock()
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
session.close.assert_awaited_once()
|
||||
yield b"chunk"
|
||||
|
||||
upstream_response = StreamingResponse(stream())
|
||||
with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)):
|
||||
response = await proxy_module.proxy(
|
||||
request, "v1/chat/completions", session=session
|
||||
)
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
chunks = [chunk async for chunk in response.body_iterator]
|
||||
assert chunks == [b"chunk"]
|
||||
@@ -1,4 +1,6 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
@@ -7,6 +9,7 @@ from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr import wallet
|
||||
from routstr.core.db import CashuTransaction
|
||||
from routstr.wallet import refund_sweep_once
|
||||
|
||||
@@ -39,6 +42,43 @@ async def _load(
|
||||
return {row.token: row for row in result.all()}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_releases_db_session_during_token_redemption(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="eligible", amount=1, unit="sat", type="out", created_at=800
|
||||
),
|
||||
)
|
||||
session_open = False
|
||||
|
||||
@asynccontextmanager
|
||||
async def tracked_session() -> AsyncIterator[AsyncSession]:
|
||||
nonlocal session_open
|
||||
async with session_factory() as session:
|
||||
session_open = True
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session_open = False
|
||||
|
||||
async def receive_token(token: str) -> None:
|
||||
assert token == "eligible"
|
||||
assert session_open is False
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", tracked_session),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=receive_token)),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
assert (await _load(session_factory))["eligible"].swept is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
@@ -90,14 +130,17 @@ async def test_refund_sweep_only_processes_expired_eligible_outgoing_tokens(
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("error", "collected"),
|
||||
("error", "collected", "claim_started_at"),
|
||||
[
|
||||
(RuntimeError("token already spent"), True),
|
||||
(RuntimeError("mint unavailable"), False),
|
||||
(RuntimeError("token already spent"), True, None),
|
||||
(RuntimeError("mint unavailable"), False, 1000),
|
||||
],
|
||||
)
|
||||
async def test_refund_sweep_records_terminal_but_not_transient_failures(
|
||||
session_factory: async_sessionmaker[AsyncSession], error: Exception, collected: bool
|
||||
async def test_refund_sweep_records_spent_and_unknown_outcomes_safely(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
error: Exception,
|
||||
collected: bool,
|
||||
claim_started_at: int | None,
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
@@ -116,3 +159,236 @@ async def test_refund_sweep_records_terminal_but_not_transient_failures(
|
||||
refund = (await _load(session_factory))["refund"]
|
||||
assert refund.collected is collected
|
||||
assert refund.swept is False
|
||||
assert refund.sweep_started_at == claim_started_at
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_spend_failure_retains_claim_and_stale_retry_records_sweep(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="post-spend-failure",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(
|
||||
side_effect=wallet.TokenConsumedError(
|
||||
"Mint on primary failed after successful melt"
|
||||
)
|
||||
),
|
||||
),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
retained = (await _load(session_factory))["post-spend-failure"]
|
||||
assert retained.swept is False
|
||||
assert retained.collected is False
|
||||
assert retained.sweep_started_at == 1000
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1300),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(side_effect=RuntimeError("token already spent")),
|
||||
),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
recovered = (await _load(session_factory))["post-spend-failure"]
|
||||
assert recovered.swept is True
|
||||
assert recovered.collected is False
|
||||
assert recovered.sweep_started_at is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_retains_claim_on_cancellation_during_redemption(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="cancelled", amount=1, unit="sat", type="out", created_at=800
|
||||
),
|
||||
)
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(side_effect=asyncio.CancelledError()),
|
||||
),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await refund_sweep_once()
|
||||
|
||||
refund = (await _load(session_factory))["cancelled"]
|
||||
assert refund.swept is False
|
||||
assert refund.sweep_started_at == 1000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_checkpoint_failure_retains_claim_and_stale_retry_records_sweep(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="checkpoint-failure",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
),
|
||||
)
|
||||
real_set_state = wallet._set_refund_sweep_state
|
||||
|
||||
async def fail_swept_checkpoint(
|
||||
refund_id: str,
|
||||
*,
|
||||
predicates: tuple[object, ...] = (),
|
||||
**values: object,
|
||||
) -> int:
|
||||
if values.get("swept") is True:
|
||||
raise RuntimeError("checkpoint unavailable")
|
||||
return await real_set_state(refund_id, predicates=predicates, **values)
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token", AsyncMock(return_value=(1, "sat", "mint"))
|
||||
),
|
||||
patch(
|
||||
"routstr.wallet._set_refund_sweep_state",
|
||||
side_effect=fail_swept_checkpoint,
|
||||
),
|
||||
patch("routstr.wallet.logger.critical") as critical,
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
retained = (await _load(session_factory))["checkpoint-failure"]
|
||||
assert retained.swept is False
|
||||
assert retained.collected is False
|
||||
assert retained.sweep_started_at == 1000
|
||||
critical.assert_called_once()
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1300),
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(side_effect=RuntimeError("token already spent")),
|
||||
),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
recovered = (await _load(session_factory))["checkpoint-failure"]
|
||||
assert recovered.swept is True
|
||||
assert recovered.collected is False
|
||||
assert recovered.sweep_started_at is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("redemption_succeeds", [True, False])
|
||||
async def test_expired_worker_cannot_overwrite_or_release_newer_claim(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
redemption_succeeds: bool,
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="reclaimed",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
),
|
||||
)
|
||||
|
||||
async def replace_claim(_token: str) -> tuple[int, str, str]:
|
||||
async with session_factory() as session:
|
||||
result = await session.exec(
|
||||
select(CashuTransaction).where(CashuTransaction.token == "reclaimed")
|
||||
)
|
||||
transaction = result.one()
|
||||
transaction.sweep_started_at = 1100
|
||||
session.add(transaction)
|
||||
await session.commit()
|
||||
if not redemption_succeeds:
|
||||
raise RuntimeError("mint unavailable")
|
||||
return (1, "sat", "mint")
|
||||
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", AsyncMock(side_effect=replace_claim)),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
reclaimed = (await _load(session_factory))["reclaimed"]
|
||||
assert reclaimed.swept is False
|
||||
assert reclaimed.collected is False
|
||||
assert reclaimed.sweep_started_at == 1100
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_sweep_recovers_stale_claim_without_misreporting_collection(
|
||||
session_factory: async_sessionmaker[AsyncSession],
|
||||
) -> None:
|
||||
await _insert(
|
||||
session_factory,
|
||||
CashuTransaction(
|
||||
token="stale",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
sweep_started_at=100,
|
||||
),
|
||||
CashuTransaction(
|
||||
token="active",
|
||||
amount=1,
|
||||
unit="sat",
|
||||
type="out",
|
||||
created_at=800,
|
||||
sweep_started_at=950,
|
||||
),
|
||||
)
|
||||
receive = AsyncMock(side_effect=RuntimeError("token already spent"))
|
||||
with (
|
||||
patch("routstr.wallet.db.create_session", side_effect=session_factory),
|
||||
patch("routstr.wallet.settings.refund_sweep_ttl_seconds", 100),
|
||||
patch("routstr.wallet.settings.refund_sweep_claim_timeout_seconds", 200),
|
||||
patch("routstr.wallet.time.time", return_value=1000),
|
||||
patch("routstr.wallet.recieve_token", receive),
|
||||
):
|
||||
await refund_sweep_once()
|
||||
|
||||
receive.assert_awaited_once_with("stale")
|
||||
loaded = await _load(session_factory)
|
||||
assert loaded["stale"].swept is True
|
||||
assert loaded["stale"].collected is False
|
||||
assert loaded["stale"].sweep_started_at is None
|
||||
assert loaded["active"].swept is False
|
||||
assert loaded["active"].sweep_started_at == 950
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
import os
|
||||
import sqlite3
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _run_alembic(root: Path, database_url: str, *args: str) -> None:
|
||||
env = os.environ.copy()
|
||||
env["DATABASE_URL"] = database_url
|
||||
subprocess.run(
|
||||
[sys.executable, "-m", "alembic", *args],
|
||||
cwd=root,
|
||||
env=env,
|
||||
check=True,
|
||||
)
|
||||
|
||||
|
||||
def test_refund_sweep_repair_restores_column_missing_at_stamped_head(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
root = Path(__file__).resolve().parents[2]
|
||||
database_path = tmp_path / "missing-refund-sweep-claim.db"
|
||||
database_url = f"sqlite+aiosqlite:///{database_path}"
|
||||
|
||||
_run_alembic(root, database_url, "upgrade", "9c4d8e2f1a6b")
|
||||
_run_alembic(root, database_url, "stamp", "aa50fde387a2")
|
||||
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
columns_before_repair = {
|
||||
row[1]
|
||||
for row in connection.execute("PRAGMA table_info(cashu_transactions)")
|
||||
}
|
||||
assert "sweep_started_at" not in columns_before_repair
|
||||
|
||||
_run_alembic(root, database_url, "upgrade", "head")
|
||||
|
||||
with sqlite3.connect(database_path) as connection:
|
||||
version = connection.execute(
|
||||
"SELECT version_num FROM alembic_version"
|
||||
).fetchone()
|
||||
columns_after_repair = {
|
||||
row[1]
|
||||
for row in connection.execute("PRAGMA table_info(cashu_transactions)")
|
||||
}
|
||||
|
||||
assert version == ("b7c9d1e3f5a2",)
|
||||
assert "sweep_started_at" in columns_after_repair
|
||||
@@ -62,6 +62,101 @@ def test_payout_settings_have_sensible_defaults() -> None:
|
||||
assert s.payout_interval_seconds == 900
|
||||
|
||||
|
||||
def test_database_pool_defaults_match_sqlalchemy_capacity() -> None:
|
||||
s = Settings()
|
||||
assert s.database_pool_size == 5
|
||||
assert s.database_max_overflow == 10
|
||||
assert s.database_pool_timeout == 30.0
|
||||
assert s.database_pool_recycle == 1800
|
||||
assert s.database_pool_pre_ping is False
|
||||
assert s.database_pool_hold_warn_seconds == 10.0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "bad_value"),
|
||||
[
|
||||
("database_pool_size", 0),
|
||||
("database_max_overflow", -1),
|
||||
("database_pool_timeout", 0),
|
||||
("database_pool_recycle", -1),
|
||||
("database_pool_hold_warn_seconds", 0),
|
||||
],
|
||||
)
|
||||
def test_database_pool_settings_reject_invalid_values(
|
||||
field: str, bad_value: int
|
||||
) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
Settings.parse_obj({field: bad_value})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_database_pool_fields_are_env_only_not_persisted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""DB pool sizing is infrastructure the node needs *before* it can read the
|
||||
DB, so it can never be configured from the DB — it must never be written to
|
||||
the settings blob, and a stale/injected DB value must never shadow env.
|
||||
"""
|
||||
monkeypatch.setenv("DATABASE_POOL_SIZE", "7")
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
s = await SettingsService.initialize(session)
|
||||
|
||||
# The env value is live for runtime consumers...
|
||||
assert s.database_pool_size == 7
|
||||
# ...but pool sizing is never written to the settings blob.
|
||||
blob = await _read_settings_blob(session)
|
||||
for field in (
|
||||
"database_pool_size",
|
||||
"database_max_overflow",
|
||||
"database_pool_timeout",
|
||||
"database_pool_recycle",
|
||||
"database_pool_pre_ping",
|
||||
"database_pool_hold_warn_seconds",
|
||||
):
|
||||
assert field not in blob
|
||||
|
||||
# Even a stale blob that somehow carries a pool value must not win: env
|
||||
# stays authoritative on the next initialize.
|
||||
await session.exec( # type: ignore
|
||||
text("UPDATE settings SET data = :d WHERE id = 1").bindparams(
|
||||
d=json.dumps({"database_pool_size": 99})
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
again = await SettingsService.initialize(session)
|
||||
assert again.database_pool_size == 7
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_does_not_apply_env_only_fields_to_live_settings(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""DB pool sizing is env-only: a settings update must neither persist it nor
|
||||
mutate the live value. The engine pool is already built at boot from env, so
|
||||
a UI/API update carrying a pool value must not make the live setting diverge
|
||||
from the running pool.
|
||||
"""
|
||||
monkeypatch.delenv("DATABASE_POOL_SIZE", raising=False)
|
||||
monkeypatch.setattr(settings, "database_pool_size", 5)
|
||||
|
||||
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
await SettingsService.initialize(session)
|
||||
await SettingsService.update(
|
||||
{"database_pool_size": 99, "name": "PoolTweaker"}, session
|
||||
)
|
||||
|
||||
# A non-env-only field still updates normally...
|
||||
assert settings.name == "PoolTweaker"
|
||||
# ...but the env-only pool size stays at the boot value.
|
||||
assert settings.database_pool_size == 5
|
||||
# ...and it is never written to the settings blob.
|
||||
blob = await _read_settings_blob(session)
|
||||
assert "database_pool_size" not in blob
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"field,bad_value",
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user