mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
Support PostgreSQL for the application database
SQLite stores every INTEGER as 64-bit, so a set of portability defects stayed invisible until the application database ran on PostgreSQL: - Monetary columns hold millisatoshis but mapped to INT4, capping a key balance at 2_147_483_647 msats (~0.0215 BTC); asyncpg rejects anything larger. Unix timestamp columns had the same width, so long-dated expiries and the 2038 boundary would overflow. Both are widened to BIGINT by a migration that is a no-op on SQLite, where the width does not exist. - Secret.nsec_state was a bare Python Enum, which makes PostgreSQL demand a native `nsecstate` type that the migration never created, so every read or write of the secrets singleton failed and took node bootstrap with it. It is now a non-native enum, rendering VARCHAR on both backends. - Model-path persistence imported the SQLite-only INSERT .. ON CONFLICT construct, which does not compile against PostgreSQL. The upsert is now built for the dialect bound to the session. - The Alembic version-clear recovery path built a sync engine from the async URL, which left postgresql+asyncpg intact and raised MissingGreenlet. It now runs through the configured async driver. - The WAL PRAGMA is gated on the bound dialect rather than a DATABASE_URL prefix. SQLite-only Cashu wallet maintenance stays scoped to .wallet/*.sqlite3 and the log-derived analytics index is untouched. - PostgreSQL migration failures and unknown revisions now fail closed instead of stamping head. PostgreSQL rolls back DDL transactionally, so stamping after an error can mark migrations applied that never committed. Existing SQLite recovery behaviour is unchanged. Adds asyncpg and two test suites: dialect-portability invariants that run in the normal suite, and an opt-in integration suite covering the full Alembic chain plus the billing, reservation, refund, invoice and payout paths against a real server, enabled by ROUTSTR_TEST_POSTGRES_URL and skipped without it.
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
"""widen money and timestamp columns to 64-bit
|
||||
|
||||
Revision ID: d3c8b21f7a04
|
||||
Revises: e4c7a1b9d520
|
||||
Create Date: 2026-09-29 00:00:00.000000
|
||||
|
||||
Balances are millisatoshis and every clock column is a unix timestamp, so both
|
||||
are 64-bit quantities. SQLite stores all INTEGER values as 64-bit, which is why
|
||||
plain ``Integer`` columns were safe there — but SQLAlchemy's ``Integer`` maps to
|
||||
PostgreSQL ``INT4``. On PostgreSQL that caps a key balance at 2_147_483_647
|
||||
msats (~0.0215 BTC, asyncpg raises "value out of int32 range" past it), lets
|
||||
lifetime counters like ``total_spent`` and ``total_paid_msats`` overflow, and
|
||||
puts every unix timestamp on the 2038 cliff — a long-dated ``key_expiry_time``
|
||||
or invoice ``validity_date`` can already exceed INT4 today.
|
||||
|
||||
ALTER TYPE is a no-op on SQLite (its INTEGER is already 64-bit and rewriting the
|
||||
tables through batch mode would needlessly churn the whole DB), so this only
|
||||
runs where the width is real.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "d3c8b21f7a04"
|
||||
down_revision = "e4c7a1b9d520"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
# (table, column, nullable) for every monetary, lifetime-counter and unix
|
||||
# timestamp column. Surrogate keys, foreign keys and ``models.context_length``
|
||||
# stay INT4: they are identifiers and bounded magnitudes, not money or clocks.
|
||||
WIDENED_COLUMNS: tuple[tuple[str, str, bool], ...] = (
|
||||
("api_keys", "balance", False),
|
||||
("api_keys", "reserved_balance", False),
|
||||
("api_keys", "total_spent", False),
|
||||
("api_keys", "total_requests", False),
|
||||
("api_keys", "reserved_at", True),
|
||||
("api_keys", "key_expiry_time", True),
|
||||
("api_keys", "created_at", True),
|
||||
("api_keys", "validity_date", True),
|
||||
("cashu_transactions", "amount", False),
|
||||
("cashu_transactions", "created_at", False),
|
||||
("cashu_transactions", "sweep_started_at", True),
|
||||
("cli_tokens", "created_at", False),
|
||||
("cli_tokens", "last_used_at", True),
|
||||
("cli_tokens", "expires_at", True),
|
||||
("lightning_invoices", "amount_sats", False),
|
||||
("lightning_invoices", "created_at", False),
|
||||
("lightning_invoices", "expires_at", False),
|
||||
("lightning_invoices", "paid_at", True),
|
||||
("lightning_invoices", "validity_date", True),
|
||||
("model_paths", "updated_at", False),
|
||||
("models", "created", False),
|
||||
("refunds", "amount_msats", False),
|
||||
("refunds", "claimed_at", True),
|
||||
("refunds", "created_at", False),
|
||||
("refunds", "updated_at", False),
|
||||
("reservation_releases", "reserved_msats", False),
|
||||
("reservation_releases", "created_at", False),
|
||||
("routstr_fees", "accumulated_msats", False),
|
||||
("routstr_fees", "total_paid_msats", False),
|
||||
("routstr_fees", "payout_in_progress_msats", False),
|
||||
("routstr_fees", "last_paid_at", True),
|
||||
("routstr_fees", "payout_started_at", True),
|
||||
("secrets", "updated_at", True),
|
||||
)
|
||||
|
||||
|
||||
def _retype(target: sa.types.TypeEngine, existing: sa.types.TypeEngine) -> None:
|
||||
bind = op.get_bind()
|
||||
if bind.dialect.name == "sqlite":
|
||||
# SQLite INTEGER is already 64-bit; a batch recreate would rewrite every
|
||||
# table for a change that does not exist on this backend.
|
||||
return
|
||||
for table, column, nullable in WIDENED_COLUMNS:
|
||||
# Missing schema is an error, not a reason to stamp a partial upgrade.
|
||||
op.alter_column(
|
||||
table,
|
||||
column,
|
||||
type_=target,
|
||||
existing_type=existing,
|
||||
existing_nullable=nullable,
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
_retype(sa.BigInteger(), sa.Integer())
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Narrowing back to INT4 fails loudly on any row that outgrew it rather than
|
||||
# silently truncating a balance, which is the correct outcome here.
|
||||
_retype(sa.Integer(), sa.BigInteger())
|
||||
@@ -8,6 +8,7 @@ requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"fastapi[standard-no-fastapi-cloud-cli]>=0.141",
|
||||
"aiosqlite>=0.20",
|
||||
"asyncpg>=0.30", # PostgreSQL backend for DATABASE_URL=postgresql+asyncpg://
|
||||
"sqlmodel>=0.0.42", # Python 3.14 deferred-annotation support
|
||||
"httpx[socks]>=0.28.1",
|
||||
"h11>=0.16",
|
||||
|
||||
+149
-38
@@ -1,18 +1,31 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import hashlib
|
||||
import os
|
||||
import pathlib
|
||||
import sqlite3
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable, Coroutine
|
||||
from contextlib import asynccontextmanager
|
||||
from enum import Enum
|
||||
from typing import AsyncGenerator
|
||||
from typing import Any, AsyncGenerator, TypeVar
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from alembic.util.exc import CommandError
|
||||
from sqlalchemy import Index, UniqueConstraint, case, delete, event, or_, text
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Column,
|
||||
Index,
|
||||
UniqueConstraint,
|
||||
case,
|
||||
delete,
|
||||
event,
|
||||
or_,
|
||||
text,
|
||||
)
|
||||
from sqlalchemy import Enum as SAEnum
|
||||
from sqlalchemy.engine import make_url
|
||||
from sqlalchemy.exc import IntegrityError, OperationalError
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine
|
||||
@@ -25,8 +38,37 @@ from .settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db")
|
||||
|
||||
# Money (millisatoshis), lifetime counters and unix timestamps are all 64-bit
|
||||
# quantities. SQLite stores every INTEGER as 64-bit so plain ``int`` fields were
|
||||
# safe there, but SQLAlchemy's ``Integer`` maps to PostgreSQL ``INT4``: a balance
|
||||
# would cap at 2_147_483_647 msats (~0.0215 BTC) and unix timestamps would break
|
||||
# in 2038. Every such column is pinned to BIGINT so both backends agree.
|
||||
Msats = BigInteger
|
||||
Sats = BigInteger
|
||||
UnixTimestamp = BigInteger
|
||||
Counter = BigInteger
|
||||
|
||||
|
||||
def _run_async_from_sync(factory: "Callable[[], Coroutine[Any, Any, T]]") -> T:
|
||||
"""Run one coroutine from sync code, in or out of a running event loop.
|
||||
|
||||
Migrations are driven synchronously by Alembic but may be triggered from
|
||||
inside FastAPI's loop, where ``asyncio.run`` would raise. Same shape as
|
||||
``migrations/env.py``: borrow a worker thread when a loop is already
|
||||
running. The coroutine is built inside the target loop — an engine created
|
||||
on one loop cannot be awaited on another.
|
||||
"""
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(factory())
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
return executor.submit(lambda: asyncio.run(factory())).result()
|
||||
|
||||
|
||||
def create_db_engine(database_url: str = DATABASE_URL) -> AsyncEngine:
|
||||
"""Build and instrument an async engine from environment-only settings."""
|
||||
@@ -97,12 +139,17 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "api_keys"
|
||||
|
||||
hashed_key: str = Field(primary_key=True)
|
||||
balance: int = Field(default=0, description="Balance in millisatoshis (msats)")
|
||||
balance: int = Field(
|
||||
default=0, sa_type=Msats, description="Balance in millisatoshis (msats)"
|
||||
)
|
||||
reserved_balance: int = Field(
|
||||
default=0, description="Reserved balance in millisatoshis (msats)"
|
||||
default=0,
|
||||
sa_type=Msats,
|
||||
description="Reserved balance in millisatoshis (msats)",
|
||||
)
|
||||
reserved_at: int | None = Field(
|
||||
default=None,
|
||||
sa_type=UnixTimestamp,
|
||||
description=(
|
||||
"Unix timestamp of the most recent balance reservation. Used to "
|
||||
"detect and release stale reservations (e.g. after client "
|
||||
@@ -115,15 +162,17 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
key_expiry_time: int | None = Field(
|
||||
default=None,
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix-timestamp after which the cashu-token's balance gets refunded to the refund_address",
|
||||
)
|
||||
total_spent: int = Field(
|
||||
default=0, description="Total spent in millisatoshis (msats)"
|
||||
default=0, sa_type=Msats, description="Total spent in millisatoshis (msats)"
|
||||
)
|
||||
total_requests: int = Field(default=0)
|
||||
total_requests: int = Field(default=0, sa_type=Counter)
|
||||
created_at: int | None = Field(
|
||||
default_factory=lambda: int(time.time()),
|
||||
nullable=True,
|
||||
sa_type=UnixTimestamp,
|
||||
description=(
|
||||
"Unix timestamp when the key was created. Nullable: keys created "
|
||||
"before this column existed have no value and sort last."
|
||||
@@ -139,6 +188,7 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
validity_date: int | None = Field(
|
||||
default=None,
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix timestamp after which the key is no longer valid",
|
||||
)
|
||||
|
||||
@@ -414,7 +464,7 @@ class ModelRow(SQLModel, table=True): # type: ignore
|
||||
primary_key=True, foreign_key="upstream_providers.id", ondelete="CASCADE"
|
||||
)
|
||||
name: str = Field()
|
||||
created: int = Field()
|
||||
created: int = Field(sa_type=UnixTimestamp)
|
||||
description: str = Field()
|
||||
context_length: int = Field()
|
||||
architecture: str = Field()
|
||||
@@ -489,6 +539,7 @@ class ModelPathRow(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
updated_at: int = Field(
|
||||
default=0,
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix timestamp of the refresh cycle that wrote this row",
|
||||
)
|
||||
|
||||
@@ -503,7 +554,7 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
|
||||
|
||||
id: str = Field(primary_key=True, description="Unique invoice identifier")
|
||||
bolt11: str = Field(description="BOLT11 invoice string", unique=True)
|
||||
amount_sats: int = Field(description="Amount in satoshis")
|
||||
amount_sats: int = Field(sa_type=Sats, description="Amount in satoshis")
|
||||
description: str = Field(description="Invoice description")
|
||||
payment_hash: str = Field(description="Payment hash for tracking", unique=True)
|
||||
status: str = Field(
|
||||
@@ -525,12 +576,19 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
|
||||
description="Mint URL where the quote was created (fallback tracking)",
|
||||
)
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), description="Unix timestamp"
|
||||
default_factory=lambda: int(time.time()),
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix timestamp",
|
||||
)
|
||||
expires_at: int = Field(
|
||||
sa_type=UnixTimestamp, description="Unix timestamp when invoice expires"
|
||||
)
|
||||
paid_at: int | None = Field(
|
||||
default=None, sa_type=UnixTimestamp, description="Unix timestamp when paid"
|
||||
)
|
||||
expires_at: int = Field(description="Unix timestamp when invoice expires")
|
||||
paid_at: int | None = Field(default=None, description="Unix timestamp when paid")
|
||||
validity_date: int | None = Field(
|
||||
default=None,
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix timestamp after which the created key expires",
|
||||
)
|
||||
|
||||
@@ -544,19 +602,21 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
|
||||
description="Unique transaction identifier",
|
||||
)
|
||||
token: str = Field(description="Serialized Cashu token")
|
||||
amount: int = Field(description="Amount in the token's unit")
|
||||
amount: int = Field(sa_type=Msats, description="Amount in the token's unit")
|
||||
unit: str = Field(description="Token unit (sat or msat)")
|
||||
mint_url: str | None = Field(default=None, description="Mint URL for the token")
|
||||
type: str = Field(default="out", description="Transaction type: in or out")
|
||||
request_id: str | None = Field(default=None, description="Associated request ID")
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()),
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix timestamp",
|
||||
)
|
||||
collected: bool = Field(default=False)
|
||||
swept: bool = Field(default=False)
|
||||
sweep_started_at: int | None = Field(
|
||||
default=None,
|
||||
sa_type=UnixTimestamp,
|
||||
description="Unix timestamp for a recoverable refund-sweep claim",
|
||||
)
|
||||
source: str = Field(
|
||||
@@ -599,7 +659,9 @@ class Refund(SQLModel, table=True): # type: ignore
|
||||
destination: str | None = Field(
|
||||
default=None, description="Lightning address or LNURL, NULL for cashu"
|
||||
)
|
||||
amount_msats: int = Field(description="Balance debited when the claim opened")
|
||||
amount_msats: int = Field(
|
||||
sa_type=Msats, description="Balance debited when the claim opened"
|
||||
)
|
||||
unit: str = Field(description="Mint unit the payout is denominated in")
|
||||
mint_url: str = Field(description="Mint the payout is drawn from")
|
||||
status: str = Field(
|
||||
@@ -612,10 +674,14 @@ class Refund(SQLModel, table=True): # type: ignore
|
||||
)
|
||||
token: str | None = Field(default=None, description="Issued cashu token")
|
||||
claimed_at: int | None = Field(
|
||||
default=None, description="Reconciler lease timestamp"
|
||||
default=None, sa_type=UnixTimestamp, description="Reconciler lease timestamp"
|
||||
)
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), sa_type=UnixTimestamp
|
||||
)
|
||||
updated_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), sa_type=UnixTimestamp
|
||||
)
|
||||
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
updated_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
|
||||
|
||||
async def store_cashu_transaction(
|
||||
@@ -783,19 +849,21 @@ class ReservationRelease(SQLModel, table=True): # type: ignore
|
||||
id: str = Field(primary_key=True)
|
||||
key_hash: str = Field(index=True)
|
||||
billing_key_hash: str = Field(index=True)
|
||||
reserved_msats: int
|
||||
reserved_msats: int = Field(sa_type=Msats)
|
||||
status: str = Field(default="active")
|
||||
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), sa_type=UnixTimestamp
|
||||
)
|
||||
|
||||
|
||||
class RoutstrFee(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "routstr_fees"
|
||||
id: int = Field(default=1, primary_key=True)
|
||||
accumulated_msats: int = Field(default=0)
|
||||
total_paid_msats: int = Field(default=0)
|
||||
last_paid_at: int | None = Field(default=None)
|
||||
payout_in_progress_msats: int = Field(default=0)
|
||||
payout_started_at: int | None = Field(default=None)
|
||||
accumulated_msats: int = Field(default=0, sa_type=Msats)
|
||||
total_paid_msats: int = Field(default=0, sa_type=Msats)
|
||||
last_paid_at: int | None = Field(default=None, sa_type=UnixTimestamp)
|
||||
payout_in_progress_msats: int = Field(default=0, sa_type=Msats)
|
||||
payout_started_at: int | None = Field(default=None, sa_type=UnixTimestamp)
|
||||
payout_quote_id: str | None = Field(default=None)
|
||||
payout_mint_url: str | None = Field(default=None)
|
||||
payout_unit: str | None = Field(default=None)
|
||||
@@ -821,6 +889,11 @@ class NsecState(str, Enum):
|
||||
cleared = "cleared"
|
||||
|
||||
|
||||
def _enum_values(enum_cls: type[Enum]) -> list[str]:
|
||||
"""Persist an enum by ``value``, matching rows written before it was typed."""
|
||||
return [str(member.value) for member in enum_cls]
|
||||
|
||||
|
||||
class Secret(SQLModel, table=True): # type: ignore
|
||||
"""Node-level secrets, stored encrypted/hashed at rest (singleton, id=1).
|
||||
|
||||
@@ -833,8 +906,21 @@ class Secret(SQLModel, table=True): # type: ignore
|
||||
id: int = Field(default=1, primary_key=True)
|
||||
admin_password_hash: str | None = Field(default=None)
|
||||
encrypted_nsec: str | None = Field(default=None)
|
||||
nsec_state: NsecState = Field(default=NsecState.legacy)
|
||||
updated_at: int | None = Field(default=None)
|
||||
# ``native_enum=False`` keeps this a VARCHAR on every backend.
|
||||
# A bare Python Enum makes SQLAlchemy reach for a native PostgreSQL ENUM
|
||||
# type named ``nsecstate``, which the migration never creates — so reading
|
||||
# or writing the secrets singleton died with `type "nsecstate" does not
|
||||
# exist`, taking node bootstrap with it. SQLite renders Enum as VARCHAR
|
||||
# either way, which is why this stayed invisible until PostgreSQL.
|
||||
nsec_state: NsecState = Field(
|
||||
default=NsecState.legacy,
|
||||
sa_column=Column(
|
||||
"nsec_state",
|
||||
SAEnum(NsecState, native_enum=False, values_callable=_enum_values),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
updated_at: int | None = Field(default=None, sa_type=UnixTimestamp)
|
||||
|
||||
|
||||
class CliToken(SQLModel, table=True): # type: ignore
|
||||
@@ -844,10 +930,14 @@ class CliToken(SQLModel, table=True): # type: ignore
|
||||
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()))
|
||||
last_used_at: int | None = Field(default=None)
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()), sa_type=UnixTimestamp
|
||||
)
|
||||
last_used_at: int | None = Field(default=None, sa_type=UnixTimestamp)
|
||||
expires_at: int | None = Field(
|
||||
default=None, description="Optional expiry unix timestamp; null = never expires"
|
||||
default=None,
|
||||
sa_type=UnixTimestamp,
|
||||
description="Optional expiry unix timestamp; null = never expires",
|
||||
)
|
||||
|
||||
|
||||
@@ -1168,7 +1258,9 @@ async def balances_by_mint_and_unit(
|
||||
async def init_db() -> None:
|
||||
"""Initializes the database and creates tables if they don't exist."""
|
||||
async with engine.begin() as conn:
|
||||
if DATABASE_URL.startswith("sqlite"):
|
||||
# WAL is a SQLite-only knob; ask the bound dialect rather than pattern
|
||||
# matching DATABASE_URL, which a caller may have swapped under us.
|
||||
if conn.dialect.name == "sqlite":
|
||||
await conn.exec_driver_sql("PRAGMA journal_mode=WAL")
|
||||
await conn.run_sync(SQLModel.metadata.create_all)
|
||||
|
||||
@@ -1226,14 +1318,24 @@ def fix_cashu_migrations() -> None:
|
||||
|
||||
|
||||
def _clear_alembic_version() -> None:
|
||||
"""Clear the alembic_version table so stamp/upgrade can proceed."""
|
||||
sync_url = DATABASE_URL.replace("+aiosqlite", "")
|
||||
from sqlalchemy import create_engine, text
|
||||
"""Clear the alembic_version table so stamp/upgrade can proceed.
|
||||
|
||||
eng = create_engine(sync_url)
|
||||
with eng.begin() as conn:
|
||||
conn.execute(text("DELETE FROM alembic_version"))
|
||||
eng.dispose()
|
||||
Driven through the configured async driver rather than a second sync engine.
|
||||
The old code built one by stripping ``+aiosqlite`` from the URL, which left
|
||||
``postgresql+asyncpg`` intact and raised ``MissingGreenlet``; naming the
|
||||
sync backend instead only moves the problem, since it then demands a
|
||||
separate sync driver (psycopg2) that a PostgreSQL deployment need not have.
|
||||
"""
|
||||
|
||||
async def clear() -> None:
|
||||
eng = create_async_engine(DATABASE_URL)
|
||||
try:
|
||||
async with eng.begin() as conn:
|
||||
await conn.execute(text("DELETE FROM alembic_version"))
|
||||
finally:
|
||||
await eng.dispose()
|
||||
|
||||
_run_async_from_sync(clear)
|
||||
|
||||
|
||||
def run_migrations() -> None:
|
||||
@@ -1260,7 +1362,13 @@ def run_migrations() -> None:
|
||||
try:
|
||||
command.upgrade(alembic_cfg, "head")
|
||||
except CommandError as e:
|
||||
if "Can't locate revision" in str(e):
|
||||
# Preserve legacy SQLite recovery only. PostgreSQL rolls back DDL
|
||||
# transactionally: stamping head after an error can hide migrations
|
||||
# that never committed. Unknown PostgreSQL revisions need operator
|
||||
# reconciliation, not an automatic stamp.
|
||||
if make_url(
|
||||
DATABASE_URL
|
||||
).get_backend_name() == "sqlite" and "Can't locate revision" in str(e):
|
||||
logger.warning(
|
||||
"Database stamped with unknown revision (likely from another branch). "
|
||||
"Re-stamping to current head.",
|
||||
@@ -1271,7 +1379,10 @@ def run_migrations() -> None:
|
||||
else:
|
||||
raise
|
||||
except OperationalError as e:
|
||||
if "duplicate column name" in str(e).lower():
|
||||
if (
|
||||
make_url(DATABASE_URL).get_backend_name() == "sqlite"
|
||||
and "duplicate column name" in str(e).lower()
|
||||
):
|
||||
logger.warning(
|
||||
"Migration hit a column that already exists (likely added via "
|
||||
"create_all on another branch). Stamping to current head.",
|
||||
|
||||
@@ -28,7 +28,8 @@ from typing import TYPE_CHECKING, Any, Callable
|
||||
from urllib.parse import parse_qsl, urlencode, urlsplit
|
||||
|
||||
import httpx
|
||||
from sqlalchemy.dialects.sqlite import insert
|
||||
from sqlalchemy.dialects.postgresql import insert as postgresql_insert
|
||||
from sqlalchemy.dialects.sqlite import insert as sqlite_insert
|
||||
from sqlalchemy.orm import selectinload
|
||||
from sqlmodel import col, delete, select
|
||||
|
||||
@@ -569,6 +570,21 @@ async def _collect_provider_paths(
|
||||
)
|
||||
|
||||
|
||||
def _upsert(session: "AsyncSession") -> Any:
|
||||
"""``INSERT .. ON CONFLICT DO UPDATE`` built for the session's own dialect.
|
||||
|
||||
``ON CONFLICT`` is spelled per-dialect in SQLAlchemy, and the SQLite
|
||||
construct does not compile against PostgreSQL — it fails at statement
|
||||
compilation with ``'OnConflictDoUpdate' object has no attribute
|
||||
'constraint_target'``, which would silently kill every model-path refresh
|
||||
(each provider's failure is caught and logged per-provider upstream).
|
||||
"""
|
||||
dialect = session.get_bind().dialect.name
|
||||
if dialect == "postgresql":
|
||||
return postgresql_insert(ModelPathRow)
|
||||
return sqlite_insert(ModelPathRow)
|
||||
|
||||
|
||||
async def _persist_provider_paths(
|
||||
upstream_provider_id: int, snapshot: ProviderPathSnapshot
|
||||
) -> None:
|
||||
@@ -602,7 +618,7 @@ async def _persist_provider_paths(
|
||||
}
|
||||
for discovered in chunk
|
||||
]
|
||||
insert_stmt = insert(ModelPathRow).values(values)
|
||||
insert_stmt = _upsert(session).values(values)
|
||||
await session.execute(
|
||||
insert_stmt.on_conflict_do_update(
|
||||
index_elements=["model_id", "path", "upstream_provider_id"],
|
||||
|
||||
@@ -63,6 +63,35 @@ USE_LOCAL_SERVICES=1 pytest tests/integration/ -v
|
||||
docker-compose -f compose.testing.yml down -v
|
||||
```
|
||||
|
||||
### PostgreSQL Mode
|
||||
|
||||
`test_postgres_compatibility.py` runs the main application DB against a real
|
||||
PostgreSQL server. It covers the full Alembic chain in both directions plus the
|
||||
billing, reservation, refund, invoice and payout paths, and it catches the
|
||||
things SQLite cannot express: native enum types, INT4 column widths, per-dialect
|
||||
`ON CONFLICT`, and genuinely concurrent transactions.
|
||||
|
||||
Every test there skips unless `ROUTSTR_TEST_POSTGRES_URL` is set, so the default
|
||||
suite is unaffected.
|
||||
|
||||
```bash
|
||||
docker run -d --name routstr-pg \
|
||||
-e POSTGRES_PASSWORD=routstr -e POSTGRES_USER=routstr -e POSTGRES_DB=routstr \
|
||||
-p 55433:5432 postgres:16-alpine
|
||||
|
||||
ROUTSTR_TEST_POSTGRES_URL=postgresql+asyncpg://routstr:routstr@127.0.0.1:55433/routstr \
|
||||
pytest tests/integration/test_postgres_compatibility.py -v
|
||||
|
||||
docker rm -f routstr-pg
|
||||
```
|
||||
|
||||
Each test drops and recreates the `public` schema, so point it only at a
|
||||
disposable database.
|
||||
|
||||
The dialect-portability invariants that do **not** need a server — column
|
||||
widths, enum rendering, per-dialect upsert selection — live in
|
||||
`tests/unit/test_postgres_schema_compat.py` and run in the normal suite.
|
||||
|
||||
### CI/CD Mode
|
||||
|
||||
```bash
|
||||
|
||||
@@ -0,0 +1,895 @@
|
||||
"""End-to-end PostgreSQL coverage for the main application DB (issue #45).
|
||||
|
||||
Runs against a real server so the things SQLite cannot express are actually
|
||||
exercised: native enum types, INT4 column widths, per-dialect ``ON CONFLICT``,
|
||||
and genuinely concurrent transactions (SQLite serialises writers, so a
|
||||
lost-update in a compare-and-swap would never show up there).
|
||||
|
||||
Point ``ROUTSTR_TEST_POSTGRES_URL`` at an empty, disposable database::
|
||||
|
||||
docker run -d --name routstr-pg -e POSTGRES_PASSWORD=routstr \\
|
||||
-e POSTGRES_USER=routstr -e POSTGRES_DB=routstr -p 55433:5432 \\
|
||||
postgres:16-alpine
|
||||
|
||||
ROUTSTR_TEST_POSTGRES_URL=postgresql+asyncpg://routstr:routstr@127.0.0.1:55433/routstr \\
|
||||
pytest tests/integration/test_postgres_compatibility.py
|
||||
|
||||
Without that variable every test here skips, so the default suite is unchanged.
|
||||
Each test gets a freshly migrated schema; the public schema is dropped between
|
||||
tests, so never aim this at a database you care about.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, AsyncIterator, Iterator
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from alembic.autogenerate import compare_metadata
|
||||
from alembic.config import Config
|
||||
from alembic.migration import MigrationContext
|
||||
from alembic.script import ScriptDirectory
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
POSTGRES_URL = os.environ.get("ROUTSTR_TEST_POSTGRES_URL", "")
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.integration,
|
||||
pytest.mark.skipif(
|
||||
not POSTGRES_URL,
|
||||
reason="set ROUTSTR_TEST_POSTGRES_URL to a disposable PostgreSQL database",
|
||||
),
|
||||
]
|
||||
|
||||
MINT = "https://mint.example"
|
||||
# 21M BTC in millisatoshis: past INT4 by nine orders of magnitude, and past
|
||||
# INT8 by none. Any monetary column that is still INT4 fails on this value.
|
||||
HUGE_MSATS = 21_000_000 * 100_000_000 * 1000
|
||||
# Comfortably past 2038-01-19, when unix seconds overflow INT4.
|
||||
POST_2038 = 2_600_000_000
|
||||
|
||||
|
||||
def _alembic(*args: str, url: str = POSTGRES_URL) -> subprocess.CompletedProcess[str]:
|
||||
"""Drive Alembic out-of-process, the way a deployment does."""
|
||||
env = os.environ.copy()
|
||||
env["DATABASE_URL"] = url
|
||||
return subprocess.run(
|
||||
[sys.executable, "-m", "alembic", *args],
|
||||
cwd=ROOT,
|
||||
env=env,
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
|
||||
def _current_revision() -> str:
|
||||
"""The stamped revision, isolated from the app's startup logging on stdout."""
|
||||
lines = [line.strip() for line in _alembic("current").stdout.splitlines()]
|
||||
revisions = [line for line in lines if line and " " not in line.rstrip(" (head)")]
|
||||
return revisions[-1].split()[0] if revisions else ""
|
||||
|
||||
|
||||
async def _reset_schema() -> None:
|
||||
engine = create_async_engine(POSTGRES_URL, poolclass=None)
|
||||
try:
|
||||
async with engine.begin() as conn:
|
||||
await conn.execute(text("DROP SCHEMA public CASCADE"))
|
||||
await conn.execute(text("CREATE SCHEMA public"))
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def migrated_postgres() -> Iterator[str]:
|
||||
"""An empty database migrated to head, torn down after the test."""
|
||||
asyncio.run(_reset_schema())
|
||||
_alembic("upgrade", "head")
|
||||
yield POSTGRES_URL
|
||||
asyncio.run(_reset_schema())
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def pg_session(migrated_postgres: str) -> AsyncIterator[AsyncSession]:
|
||||
engine = create_async_engine(migrated_postgres)
|
||||
try:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
yield session
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _new_key(session: AsyncSession, hashed_key: str, **kwargs: Any) -> Any:
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
key = ApiKey(
|
||||
hashed_key=hashed_key,
|
||||
refund_mint_url=MINT,
|
||||
refund_currency="sat",
|
||||
**kwargs,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
return key
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Alembic chain
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_full_migration_chain_upgrades_and_downgrades(migrated_postgres: str) -> None:
|
||||
"""Every revision must run both ways on PostgreSQL, not just on SQLite."""
|
||||
head = ScriptDirectory.from_config(
|
||||
Config(str(ROOT / "alembic.ini"))
|
||||
).get_current_head()
|
||||
assert _current_revision() == head
|
||||
|
||||
_alembic("downgrade", "base")
|
||||
assert _current_revision() == ""
|
||||
|
||||
_alembic("upgrade", "head")
|
||||
assert _current_revision() == head
|
||||
|
||||
|
||||
def test_migrated_schema_matches_the_orm(migrated_postgres: str) -> None:
|
||||
"""A migrated database and ``SQLModel.metadata`` must not disagree.
|
||||
|
||||
Drift here is how ``secrets.nsec_state`` shipped: the migration made it
|
||||
VARCHAR while the ORM expected a native ``nsecstate`` enum type that
|
||||
nothing ever created.
|
||||
"""
|
||||
import routstr.core.db # noqa: F401 - registers every table
|
||||
|
||||
async def compare() -> list[Any]:
|
||||
engine = create_async_engine(migrated_postgres)
|
||||
try:
|
||||
async with engine.connect() as conn:
|
||||
return await conn.run_sync(
|
||||
lambda sync_conn: compare_metadata(
|
||||
MigrationContext.configure(
|
||||
sync_conn, opts={"compare_type": True}
|
||||
),
|
||||
SQLModel.metadata,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
diffs = asyncio.run(compare())
|
||||
|
||||
def _touches(diff: Any, table: str, column: str | None = None) -> bool:
|
||||
rendered = str(diff)
|
||||
return table in rendered and (column is None or column in rendered)
|
||||
|
||||
# Known, pre-existing and backend-independent (identical on SQLite):
|
||||
# - `settings` is managed by raw SQL, so it is absent from the ORM metadata
|
||||
# - `cashu_transactions.api_key_hashed_key` declares an FK neither backend has
|
||||
# - `cli_tokens.token` carries both a unique constraint and a unique index
|
||||
# - TEXT vs VARCHAR is not a behavioural difference in PostgreSQL
|
||||
unexplained = [
|
||||
diff
|
||||
for diff in diffs
|
||||
if not (
|
||||
_touches(diff, "settings")
|
||||
or _touches(diff, "cashu_transactions", "api_key_hashed_key")
|
||||
or _touches(diff, "cli_tokens", "token")
|
||||
or "modify_type" in str(diff)
|
||||
and "TEXT()" in str(diff)
|
||||
)
|
||||
]
|
||||
assert unexplained == [], unexplained
|
||||
|
||||
|
||||
def test_no_column_that_holds_money_or_time_is_int4(migrated_postgres: str) -> None:
|
||||
"""INT4 msats caps a balance at ~0.0215 BTC; INT4 unix seconds die in 2038."""
|
||||
|
||||
async def widths() -> list[str]:
|
||||
engine = create_async_engine(POSTGRES_URL)
|
||||
try:
|
||||
async with engine.connect() as conn:
|
||||
result = await conn.execute(
|
||||
text(
|
||||
"SELECT table_name, column_name FROM information_schema.columns"
|
||||
" WHERE table_schema = 'public' AND data_type = 'integer'"
|
||||
)
|
||||
)
|
||||
return [f"{row[0]}.{row[1]}" for row in result]
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
narrow = set(asyncio.run(widths()))
|
||||
# Surrogate keys, foreign keys and a bounded model attribute.
|
||||
allowed = {
|
||||
"model_paths.id",
|
||||
"model_paths.upstream_provider_id",
|
||||
"models.context_length",
|
||||
"models.upstream_provider_id",
|
||||
"routstr_fees.id",
|
||||
"secrets.id",
|
||||
"settings.id",
|
||||
"upstream_providers.id",
|
||||
}
|
||||
assert narrow <= allowed, f"still INT4: {sorted(narrow - allowed)}"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Billing
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_holds_more_than_int4(pg_session: AsyncSession) -> None:
|
||||
"""A node holding more than ~0.0215 BTC must not fail to write its balance."""
|
||||
from routstr.core.db import ApiKey, total_user_liability
|
||||
|
||||
await _new_key(pg_session, "whale", balance=HUGE_MSATS, total_spent=HUGE_MSATS)
|
||||
|
||||
stored = await pg_session.get(ApiKey, "whale")
|
||||
assert stored is not None
|
||||
assert stored.balance == HUGE_MSATS
|
||||
assert stored.total_spent == HUGE_MSATS
|
||||
assert await total_user_liability(pg_session) == HUGE_MSATS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_long_dated_expiries_round_trip(pg_session: AsyncSession) -> None:
|
||||
"""Expiry timestamps past 2038 must survive the round trip."""
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
await _new_key(
|
||||
pg_session,
|
||||
"long-lived",
|
||||
balance=1_000,
|
||||
key_expiry_time=POST_2038,
|
||||
validity_date=POST_2038,
|
||||
)
|
||||
stored = await pg_session.get(ApiKey, "long-lived")
|
||||
assert stored is not None
|
||||
assert stored.key_expiry_time == POST_2038
|
||||
assert stored.validity_date == POST_2038
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_aggregates_across_mint_and_unit(
|
||||
pg_session: AsyncSession,
|
||||
) -> None:
|
||||
from routstr.core.db import (
|
||||
balance_for_mint_and_unit,
|
||||
balances_by_mint_and_unit,
|
||||
user_liability_for_mint_and_unit,
|
||||
)
|
||||
|
||||
await _new_key(pg_session, "a", balance=3_000)
|
||||
await _new_key(pg_session, "b", balance=4_500)
|
||||
|
||||
assert await balance_for_mint_and_unit(pg_session, MINT, "sat") == 7_500
|
||||
assert await user_liability_for_mint_and_unit(pg_session, MINT, "sat") == 7_500
|
||||
assert await balances_by_mint_and_unit(pg_session, [MINT], ["sat"]) == {
|
||||
(MINT, "sat"): 7_500
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_reservations_cannot_overspend(migrated_postgres: str) -> None:
|
||||
"""The reservation compare-and-swap must hold under real concurrency.
|
||||
|
||||
PostgreSQL runs these two transactions at the same time; SQLite would have
|
||||
serialised them and hidden a lost update.
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from routstr.auth import pay_for_request
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
engine = create_async_engine(migrated_postgres)
|
||||
try:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as setup:
|
||||
await _new_key(setup, "contended", balance=1_000)
|
||||
|
||||
barrier = asyncio.Barrier(2)
|
||||
|
||||
async def reserve(amount: int) -> int:
|
||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||
key = await session.get(ApiKey, "contended")
|
||||
assert key is not None
|
||||
await barrier.wait() # Both read the same pre-reservation balance.
|
||||
try:
|
||||
await pay_for_request(key, amount, session)
|
||||
except HTTPException as exc:
|
||||
assert exc.status_code == 402
|
||||
return 0
|
||||
return 1
|
||||
|
||||
# Both want 600 of a 1000 balance; exactly one may win.
|
||||
won = await asyncio.gather(reserve(600), reserve(600))
|
||||
assert sum(won) == 1
|
||||
|
||||
async with AsyncSession(engine, expire_on_commit=False) as check:
|
||||
key = await check.get(ApiKey, "contended")
|
||||
assert key is not None
|
||||
assert key.reserved_balance == 600
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_real_billing_settlement_and_replay(
|
||||
pg_session: AsyncSession, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
||||
from routstr.payment.cost_calculation import CostData
|
||||
|
||||
key = await _new_key(pg_session, "billing", balance=5_000_000_000)
|
||||
reservation = await pay_for_request(key, 3_000_000_000, pg_session)
|
||||
cost = CostData(
|
||||
base_msats=0,
|
||||
input_msats=1_000_000_000,
|
||||
output_msats=1_500_000_000,
|
||||
total_msats=2_500_000_000,
|
||||
total_usd=0.0,
|
||||
input_tokens=50,
|
||||
output_tokens=50,
|
||||
)
|
||||
|
||||
async def calculated_cost(*args: Any, **kwargs: Any) -> CostData:
|
||||
return cost
|
||||
|
||||
monkeypatch.setattr("routstr.auth.calculate_cost", calculated_cost)
|
||||
for _ in range(2):
|
||||
await adjust_payment_for_tokens(
|
||||
key,
|
||||
{
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 50, "completion_tokens": 50},
|
||||
},
|
||||
pg_session,
|
||||
reservation.reserved_msats,
|
||||
None,
|
||||
None,
|
||||
reservation,
|
||||
)
|
||||
await pg_session.refresh(key)
|
||||
assert key.balance == 2_500_000_000
|
||||
assert key.total_spent == 2_500_000_000
|
||||
assert key.reserved_balance == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_release_and_settle_are_idempotent(
|
||||
pg_session: AsyncSession,
|
||||
) -> None:
|
||||
from routstr.refund import open_claim, release, settle
|
||||
|
||||
key = await _new_key(pg_session, "claim", balance=5_000_000_000)
|
||||
claim = await open_claim(pg_session, key, method="cashu", destination=None)
|
||||
await pg_session.refresh(key)
|
||||
assert key.balance == 0 and claim.amount_msats == 5_000_000_000
|
||||
assert await release(pg_session, claim)
|
||||
assert not await release(pg_session, claim)
|
||||
await pg_session.refresh(key)
|
||||
assert key.balance == 5_000_000_000
|
||||
second = await open_claim(pg_session, key, method="cashu", destination=None)
|
||||
assert await settle(pg_session, second, token="test-token")
|
||||
assert not await settle(pg_session, second, token="test-token")
|
||||
assert not await release(pg_session, second)
|
||||
await pg_session.refresh(key)
|
||||
assert key.balance == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("purpose", ["create", "topup"])
|
||||
async def test_invoice_credit_is_atomic_and_idempotent(
|
||||
pg_session: AsyncSession, purpose: str
|
||||
) -> None:
|
||||
from routstr.core.db import ApiKey, LightningInvoice
|
||||
from routstr.lightning import _finalize_invoice_settlement, _InvoiceSettlement
|
||||
|
||||
if purpose == "topup":
|
||||
await _new_key(pg_session, "credit", balance=3_000_000_000)
|
||||
invoice = LightningInvoice(
|
||||
id="credit-invoice",
|
||||
payment_hash="credit-hash",
|
||||
bolt11="lnbc-credit",
|
||||
amount_sats=1_000_000,
|
||||
description="credit",
|
||||
purpose=purpose,
|
||||
api_key_hash="credit" if purpose == "topup" else None,
|
||||
expires_at=POST_2038,
|
||||
validity_date=POST_2038,
|
||||
mint_url=MINT,
|
||||
)
|
||||
pg_session.add(invoice)
|
||||
await pg_session.commit()
|
||||
snapshot = _InvoiceSettlement.from_invoice(invoice)
|
||||
paid, key_hash = await _finalize_invoice_settlement(snapshot, pg_session, POST_2038)
|
||||
assert paid and key_hash
|
||||
assert await _finalize_invoice_settlement(snapshot, pg_session, POST_2038) == (
|
||||
False,
|
||||
None,
|
||||
)
|
||||
key = await pg_session.get(ApiKey, key_hash)
|
||||
assert key is not None
|
||||
assert key.balance == (4_000_000_000 if purpose == "topup" else 1_000_000_000)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Reservations
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_reservations_are_released(pg_session: AsyncSession) -> None:
|
||||
from routstr.core.db import ApiKey, ReservationRelease, release_stale_reservations
|
||||
|
||||
now = int(time.time())
|
||||
await _new_key(pg_session, "resv", balance=10_000, reserved_balance=2_500)
|
||||
pg_session.add(
|
||||
ReservationRelease(
|
||||
id=uuid.uuid4().hex,
|
||||
key_hash="resv",
|
||||
billing_key_hash="resv",
|
||||
reserved_msats=2_500,
|
||||
status="active",
|
||||
created_at=now - 3_600,
|
||||
)
|
||||
)
|
||||
await pg_session.commit()
|
||||
|
||||
assert await release_stale_reservations(pg_session, max_age_seconds=60) == 1
|
||||
|
||||
released = (await pg_session.exec(select(ReservationRelease))).first()
|
||||
assert released is not None
|
||||
assert released.status == "released"
|
||||
key = await pg_session.get(ApiKey, "resv")
|
||||
assert key is not None and key.reserved_balance == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fresh_reservations_survive_the_sweep(pg_session: AsyncSession) -> None:
|
||||
from routstr.core.db import ReservationRelease, release_stale_reservations
|
||||
|
||||
await _new_key(pg_session, "fresh", balance=10_000, reserved_balance=1_000)
|
||||
pg_session.add(
|
||||
ReservationRelease(
|
||||
id=uuid.uuid4().hex,
|
||||
key_hash="fresh",
|
||||
billing_key_hash="fresh",
|
||||
reserved_msats=1_000,
|
||||
status="active",
|
||||
created_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
await pg_session.commit()
|
||||
|
||||
assert await release_stale_reservations(pg_session, max_age_seconds=3_600) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reset_all_reserved_balances(pg_session: AsyncSession) -> None:
|
||||
from routstr.core.db import ApiKey, reset_all_reserved_balances
|
||||
|
||||
await _new_key(pg_session, "r1", balance=5_000, reserved_balance=500)
|
||||
await reset_all_reserved_balances(pg_session)
|
||||
|
||||
key = await pg_session.get(ApiKey, "r1")
|
||||
assert key is not None
|
||||
assert key.reserved_balance == 0
|
||||
assert key.reserved_at is None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Refunds
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_one_open_refund_per_key(pg_session: AsyncSession) -> None:
|
||||
"""The partial unique index must be enforced by PostgreSQL too.
|
||||
|
||||
It is declared with both ``sqlite_where`` and ``postgresql_where``; if the
|
||||
PostgreSQL variant were dropped, a key could open two concurrent payouts.
|
||||
"""
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
from routstr.core.db import Refund
|
||||
|
||||
await _new_key(pg_session, "refundee", balance=50_000)
|
||||
|
||||
def _refund(status: str) -> Refund:
|
||||
return Refund(
|
||||
id=uuid.uuid4().hex,
|
||||
api_key_hashed_key="refundee",
|
||||
method="lightning",
|
||||
amount_msats=10_000,
|
||||
unit="sat",
|
||||
mint_url=MINT,
|
||||
status=status,
|
||||
)
|
||||
|
||||
pg_session.add(_refund("pending"))
|
||||
await pg_session.commit()
|
||||
|
||||
pg_session.add(_refund("ambiguous"))
|
||||
with pytest.raises(IntegrityError):
|
||||
await pg_session.commit()
|
||||
await pg_session.rollback()
|
||||
|
||||
# A closed claim does not occupy the slot.
|
||||
pg_session.add(_refund("paid"))
|
||||
await pg_session.commit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_amount_holds_more_than_int4(pg_session: AsyncSession) -> None:
|
||||
from routstr.core.db import Refund
|
||||
|
||||
await _new_key(pg_session, "big-refund", balance=HUGE_MSATS)
|
||||
refund = Refund(
|
||||
id=uuid.uuid4().hex,
|
||||
api_key_hashed_key="big-refund",
|
||||
method="cashu",
|
||||
amount_msats=HUGE_MSATS,
|
||||
unit="sat",
|
||||
mint_url=MINT,
|
||||
status="pending",
|
||||
claimed_at=POST_2038,
|
||||
)
|
||||
pg_session.add(refund)
|
||||
await pg_session.commit()
|
||||
|
||||
stored = await pg_session.get(Refund, refund.id)
|
||||
assert stored is not None
|
||||
assert stored.amount_msats == HUGE_MSATS
|
||||
assert stored.claimed_at == POST_2038
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_total_liability_counts_unresolved_refunds(
|
||||
pg_session: AsyncSession,
|
||||
) -> None:
|
||||
from routstr.core.db import Refund, total_user_liability
|
||||
|
||||
await _new_key(pg_session, "liable", balance=1_000)
|
||||
pg_session.add(
|
||||
Refund(
|
||||
id=uuid.uuid4().hex,
|
||||
api_key_hashed_key="liable",
|
||||
method="lightning",
|
||||
amount_msats=250,
|
||||
unit="sat",
|
||||
mint_url=MINT,
|
||||
status="pending",
|
||||
)
|
||||
)
|
||||
await pg_session.commit()
|
||||
|
||||
assert await total_user_liability(pg_session) == 1_250
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Invoices and payouts
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lightning_payout_lifecycle(pg_session: AsyncSession) -> None:
|
||||
from routstr.core.db import (
|
||||
LightningInvoice,
|
||||
list_unsettled_lightning_payouts,
|
||||
record_lightning_payout,
|
||||
settle_lightning_payout,
|
||||
)
|
||||
|
||||
await record_lightning_payout(
|
||||
pg_session,
|
||||
quote_id="quote-1",
|
||||
bolt11="lnbc1payout",
|
||||
amount_sats=1_234,
|
||||
mint_url=MINT,
|
||||
destination="node@example.com",
|
||||
)
|
||||
|
||||
unsettled = await list_unsettled_lightning_payouts(
|
||||
pg_session, MINT, created_before=int(time.time()) + 60
|
||||
)
|
||||
assert [row.payment_hash for row in unsettled] == ["quote-1"]
|
||||
|
||||
await settle_lightning_payout(
|
||||
pg_session, "quote-1", status="paid", amount_sats=1_234
|
||||
)
|
||||
assert (
|
||||
await list_unsettled_lightning_payouts(
|
||||
pg_session, MINT, created_before=int(time.time()) + 60
|
||||
)
|
||||
== []
|
||||
)
|
||||
|
||||
paid = (
|
||||
await pg_session.exec(
|
||||
select(LightningInvoice).where(LightningInvoice.payment_hash == "quote-1")
|
||||
)
|
||||
).one()
|
||||
assert paid.status == "paid"
|
||||
assert paid.paid_at is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invoice_unique_constraints_hold(pg_session: AsyncSession) -> None:
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
from routstr.core.db import LightningInvoice
|
||||
|
||||
def _invoice(invoice_id: str) -> LightningInvoice:
|
||||
return LightningInvoice(
|
||||
id=invoice_id,
|
||||
bolt11="lnbc1duplicate",
|
||||
amount_sats=10,
|
||||
description="dup",
|
||||
payment_hash="hash-dup",
|
||||
purpose="topup",
|
||||
expires_at=POST_2038,
|
||||
)
|
||||
|
||||
pg_session.add(_invoice("inv-1"))
|
||||
await pg_session.commit()
|
||||
|
||||
pg_session.add(_invoice("inv-2"))
|
||||
with pytest.raises(IntegrityError):
|
||||
await pg_session.commit()
|
||||
await pg_session.rollback()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_payout_checkpoint_round_trip(pg_session: AsyncSession) -> None:
|
||||
"""Accumulate, checkpoint, restore, checkpoint again, complete."""
|
||||
from routstr.core.db import (
|
||||
accumulate_routstr_fee,
|
||||
complete_routstr_fee_payout,
|
||||
get_routstr_fee,
|
||||
reset_routstr_fee,
|
||||
restore_routstr_fee_payout,
|
||||
)
|
||||
|
||||
await accumulate_routstr_fee(pg_session, HUGE_MSATS)
|
||||
assert (await get_routstr_fee(pg_session)).accumulated_msats == HUGE_MSATS
|
||||
|
||||
assert await reset_routstr_fee(pg_session, 5_000, "q1", MINT, "sat") is True
|
||||
fee = await get_routstr_fee(pg_session)
|
||||
assert fee.payout_in_progress_msats == 5_000
|
||||
assert fee.accumulated_msats == HUGE_MSATS - 5_000
|
||||
|
||||
# A payout that never landed goes back to the balance.
|
||||
assert (
|
||||
await restore_routstr_fee_payout(pg_session, 5_000, "q1", MINT, "sat") is True
|
||||
)
|
||||
fee = await get_routstr_fee(pg_session)
|
||||
assert fee.payout_in_progress_msats == 0
|
||||
assert fee.accumulated_msats == HUGE_MSATS
|
||||
|
||||
assert await reset_routstr_fee(pg_session, 5_000, "q2", MINT, "sat") is True
|
||||
assert (
|
||||
await complete_routstr_fee_payout(pg_session, 5_000, "q2", MINT, "sat") is True
|
||||
)
|
||||
fee = await get_routstr_fee(pg_session)
|
||||
assert fee.payout_in_progress_msats == 0
|
||||
assert fee.total_paid_msats == 5_000
|
||||
assert fee.accumulated_msats == HUGE_MSATS - 5_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fee_checkpoint_rejects_a_mismatched_quote(
|
||||
pg_session: AsyncSession,
|
||||
) -> None:
|
||||
from routstr.core.db import (
|
||||
accumulate_routstr_fee,
|
||||
complete_routstr_fee_payout,
|
||||
reset_routstr_fee,
|
||||
)
|
||||
|
||||
await accumulate_routstr_fee(pg_session, 10_000)
|
||||
assert await reset_routstr_fee(pg_session, 4_000, "real", MINT, "sat") is True
|
||||
assert (
|
||||
await complete_routstr_fee_payout(pg_session, 4_000, "wrong", MINT, "sat")
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------
|
||||
# Secrets, transactions and model paths
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_secrets_singleton_round_trips(pg_session: AsyncSession) -> None:
|
||||
"""``nsec_state`` used to need a native ``nsecstate`` type nothing created."""
|
||||
from routstr.core.db import NsecState, Secret, get_secret, set_admin_password
|
||||
|
||||
secret = await get_secret(pg_session)
|
||||
assert secret.nsec_state == NsecState.legacy
|
||||
|
||||
await set_admin_password(pg_session, "correct-horse-battery")
|
||||
stored = await pg_session.get(Secret, 1)
|
||||
assert stored is not None
|
||||
assert stored.admin_password_hash
|
||||
|
||||
stored.nsec_state = NsecState.cleared
|
||||
pg_session.add(stored)
|
||||
await pg_session.commit()
|
||||
|
||||
reread = await get_secret(pg_session)
|
||||
assert reread.nsec_state == NsecState.cleared
|
||||
|
||||
# Stored by value, so rows written before the column was typed still load.
|
||||
raw = await pg_session.exec(text("SELECT nsec_state FROM secrets WHERE id = 1")) # type: ignore[call-overload]
|
||||
assert raw.one()[0] == "cleared"
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def db_bound_to_postgres(
|
||||
migrated_postgres: str, monkeypatch: pytest.MonkeyPatch
|
||||
) -> AsyncIterator[Any]:
|
||||
"""Point ``routstr.core.db``'s module-level engine/URL at the test database.
|
||||
|
||||
Reloading the module is not an option: re-executing it redefines every table
|
||||
on the shared ``SQLModel.metadata``. Rebinding the two globals that the
|
||||
functions under test read is enough and leaves the mappers alone.
|
||||
"""
|
||||
from routstr.core import db as db_module
|
||||
|
||||
engine = db_module.create_db_engine(migrated_postgres)
|
||||
monkeypatch.setattr(db_module, "engine", engine)
|
||||
monkeypatch.setattr(db_module, "DATABASE_URL", migrated_postgres)
|
||||
monkeypatch.setenv("DATABASE_URL", migrated_postgres)
|
||||
try:
|
||||
yield db_module
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cashu_transaction_write_is_idempotent(db_bound_to_postgres: Any) -> None:
|
||||
"""The retry path reuses a deterministic id, so a replay must not duplicate."""
|
||||
db_module = db_bound_to_postgres
|
||||
|
||||
for _ in range(2):
|
||||
assert (
|
||||
await db_module.store_cashu_transaction_with_retry(
|
||||
token="cashuAtoken",
|
||||
amount=HUGE_MSATS,
|
||||
unit="msat",
|
||||
mint_url=MINT,
|
||||
typ="in",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
async with db_module.create_session() as session:
|
||||
rows = (await session.exec(select(db_module.CashuTransaction))).all()
|
||||
assert len(rows) == 1
|
||||
assert rows[0].amount == HUGE_MSATS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_path_upsert_updates_on_conflict(
|
||||
pg_session: AsyncSession,
|
||||
) -> None:
|
||||
"""``ON CONFLICT`` is per-dialect; the SQLite construct cannot compile here."""
|
||||
from routstr.core.db import ModelPathRow, UpstreamProviderRow
|
||||
from routstr.upstream.model_paths import _upsert
|
||||
|
||||
provider = UpstreamProviderRow(
|
||||
provider_type="custom",
|
||||
base_url="https://upstream.example",
|
||||
api_key="k",
|
||||
enabled=True,
|
||||
)
|
||||
pg_session.add(provider)
|
||||
await pg_session.commit()
|
||||
await pg_session.refresh(provider)
|
||||
|
||||
async def write(slug: str, updated_at: int) -> None:
|
||||
statement = _upsert(pg_session).values(
|
||||
[
|
||||
{
|
||||
"model_id": "gpt-x",
|
||||
"path": "url=https%3A%2F%2Fupstream.example",
|
||||
"provider_slug": slug,
|
||||
"provider_type": "custom",
|
||||
"endpoint_tag": None,
|
||||
"endpoint_name": None,
|
||||
"model_metadata": json.dumps({"slug": slug}),
|
||||
"upstream_provider_id": provider.id,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
]
|
||||
)
|
||||
await pg_session.execute(
|
||||
statement.on_conflict_do_update(
|
||||
index_elements=["model_id", "path", "upstream_provider_id"],
|
||||
set_={
|
||||
"provider_slug": statement.excluded.provider_slug,
|
||||
"model_metadata": statement.excluded.model_metadata,
|
||||
"updated_at": statement.excluded.updated_at,
|
||||
},
|
||||
)
|
||||
)
|
||||
await pg_session.commit()
|
||||
|
||||
await write("first", 1_000)
|
||||
await write("second", POST_2038)
|
||||
|
||||
rows = (await pg_session.exec(select(ModelPathRow))).all()
|
||||
assert len(rows) == 1
|
||||
assert rows[0].provider_slug == "second"
|
||||
assert rows[0].updated_at == POST_2038
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dead_keys_are_pruned(pg_session: AsyncSession) -> None:
|
||||
from routstr.core.db import ApiKey, prune_dead_api_keys
|
||||
|
||||
await _new_key(pg_session, "dead", balance=0, created_at=int(time.time()) - 10_000)
|
||||
await _new_key(pg_session, "alive", balance=1_000)
|
||||
|
||||
assert await prune_dead_api_keys(pg_session, min_age_seconds=60) == 1
|
||||
assert await pg_session.get(ApiKey, "dead") is None
|
||||
assert await pg_session.get(ApiKey, "alive") is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_version_clear_uses_async_driver(db_bound_to_postgres: Any) -> None:
|
||||
db_module = db_bound_to_postgres
|
||||
# Direct invocation really exercises the running-loop branch.
|
||||
db_module._clear_alembic_version()
|
||||
async with db_module.engine.connect() as conn:
|
||||
assert (
|
||||
await conn.execute(text("SELECT count(*) FROM alembic_version"))
|
||||
).scalar_one() == 0
|
||||
# Explicit stamp is justified only here: this fixture is known to be at head.
|
||||
await asyncio.to_thread(_alembic, "stamp", "head")
|
||||
await asyncio.to_thread(db_module._clear_alembic_version)
|
||||
async with db_module.engine.connect() as conn:
|
||||
assert (
|
||||
await conn.execute(text("SELECT count(*) FROM alembic_version"))
|
||||
).scalar_one() == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("unknown_revision", [False, True])
|
||||
async def test_postgres_migration_errors_never_stamp_head(
|
||||
db_bound_to_postgres: Any, unknown_revision: bool
|
||||
) -> None:
|
||||
from alembic.util.exc import CommandError
|
||||
from sqlalchemy.exc import ProgrammingError
|
||||
|
||||
db_module = db_bound_to_postgres
|
||||
async with db_module.engine.begin() as conn:
|
||||
if unknown_revision:
|
||||
await conn.execute(
|
||||
text("UPDATE alembic_version SET version_num = 'unknown_revision'")
|
||||
)
|
||||
else:
|
||||
# An empty version table and existing tables triggers duplicate DDL.
|
||||
await conn.execute(text("DELETE FROM alembic_version"))
|
||||
with pytest.raises((CommandError, ProgrammingError)):
|
||||
await asyncio.to_thread(db_module.run_migrations)
|
||||
async with db_module.engine.connect() as conn:
|
||||
versions = (
|
||||
(await conn.execute(text("SELECT version_num FROM alembic_version")))
|
||||
.scalars()
|
||||
.all()
|
||||
)
|
||||
assert versions == (["unknown_revision"] if unknown_revision else [])
|
||||
@@ -0,0 +1,251 @@
|
||||
"""Dialect-portability invariants for the main application DB.
|
||||
|
||||
These run without a PostgreSQL server: they compile the real SQLModel metadata
|
||||
against the PostgreSQL dialect and assert the properties that SQLite silently
|
||||
papered over. Each test here corresponds to a defect that was live on
|
||||
PostgreSQL while every SQLite test stayed green (issue #45).
|
||||
|
||||
The end-to-end proof against a real server lives in
|
||||
``tests/integration/test_postgres_compatibility.py``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import BigInteger, Integer
|
||||
from sqlalchemy.dialects import postgresql, sqlite
|
||||
from sqlalchemy.schema import CreateTable
|
||||
from sqlmodel import SQLModel
|
||||
|
||||
from routstr.core import db as db_module
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
MIGRATION = (
|
||||
ROOT
|
||||
/ "migrations"
|
||||
/ "versions"
|
||||
/ "d3c8b21f7a04_widen_money_and_timestamp_columns.py"
|
||||
)
|
||||
|
||||
# Built once here because SQLAlchemy's dialect constructors are untyped.
|
||||
POSTGRES_DIALECT: Any = postgresql.dialect() # type: ignore[no-untyped-call]
|
||||
SQLITE_DIALECT: Any = sqlite.dialect() # type: ignore[no-untyped-call]
|
||||
|
||||
|
||||
def _load_widening_migration() -> Any:
|
||||
spec = importlib.util.spec_from_file_location("_widen_migration", MIGRATION)
|
||||
assert spec is not None and spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _table(name: str) -> Any:
|
||||
return SQLModel.metadata.tables[name]
|
||||
|
||||
|
||||
def _render(element: Any, dialect: Any) -> str:
|
||||
"""Compile one element against a dialect and return the rendered SQL."""
|
||||
return str(element.compile(dialect=dialect))
|
||||
|
||||
|
||||
def test_migration_and_models_widen_the_same_columns() -> None:
|
||||
"""The ALTER list and the ORM must not drift apart.
|
||||
|
||||
If they do, a fresh PostgreSQL database (built by the migration) and the
|
||||
in-memory metadata (used by ``create_all``) disagree on column width, and
|
||||
only one of the two paths overflows.
|
||||
"""
|
||||
widened = _load_widening_migration().WIDENED_COLUMNS
|
||||
|
||||
mismatched = []
|
||||
for table_name, column_name, nullable in widened:
|
||||
column = _table(table_name).columns[column_name]
|
||||
if not isinstance(column.type, BigInteger):
|
||||
mismatched.append(
|
||||
f"{table_name}.{column_name}: migration widens it, model has "
|
||||
f"{column.type!r}"
|
||||
)
|
||||
if column.nullable != nullable:
|
||||
mismatched.append(
|
||||
f"{table_name}.{column_name} nullable={column.nullable}, "
|
||||
f"migration says {nullable}"
|
||||
)
|
||||
assert mismatched == []
|
||||
|
||||
# And the other direction, so a column widened in the ORM but left out of
|
||||
# the ALTER list cannot pass: ``create_all`` would give it BIGINT while a
|
||||
# migrated database kept INT4.
|
||||
declared = {(table, column) for table, column, _ in widened}
|
||||
in_models = {
|
||||
(table.name, column.name)
|
||||
for table in SQLModel.metadata.tables.values()
|
||||
for column in table.columns
|
||||
if isinstance(column.type, BigInteger)
|
||||
}
|
||||
assert in_models - declared == set(), "widened in models but not in migration"
|
||||
assert declared - in_models == set(), "widened in migration but not in models"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"table_name, column_name",
|
||||
[
|
||||
("api_keys", "balance"),
|
||||
("api_keys", "reserved_balance"),
|
||||
("api_keys", "total_spent"),
|
||||
("api_keys", "total_requests"),
|
||||
("cashu_transactions", "amount"),
|
||||
("lightning_invoices", "amount_sats"),
|
||||
("refunds", "amount_msats"),
|
||||
("reservation_releases", "reserved_msats"),
|
||||
("routstr_fees", "accumulated_msats"),
|
||||
("routstr_fees", "payout_in_progress_msats"),
|
||||
("routstr_fees", "total_paid_msats"),
|
||||
],
|
||||
)
|
||||
def test_money_columns_are_64_bit_on_postgresql(
|
||||
table_name: str, column_name: str
|
||||
) -> None:
|
||||
"""Balances are millisatoshis; INT4 would cap a key at ~0.0215 BTC.
|
||||
|
||||
SQLite stores every INTEGER as 64-bit, so a plain ``int`` field was safe
|
||||
there and this ceiling existed only on PostgreSQL.
|
||||
"""
|
||||
column = _table(table_name).columns[column_name]
|
||||
rendered = _render(column.type, POSTGRES_DIALECT)
|
||||
assert rendered == "BIGINT", f"{table_name}.{column_name} renders as {rendered}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"table_name, column_name",
|
||||
[
|
||||
("api_keys", "created_at"),
|
||||
("api_keys", "key_expiry_time"),
|
||||
("api_keys", "reserved_at"),
|
||||
("api_keys", "validity_date"),
|
||||
("cashu_transactions", "created_at"),
|
||||
("cli_tokens", "expires_at"),
|
||||
("lightning_invoices", "expires_at"),
|
||||
("lightning_invoices", "paid_at"),
|
||||
("refunds", "claimed_at"),
|
||||
("reservation_releases", "created_at"),
|
||||
("routstr_fees", "payout_started_at"),
|
||||
],
|
||||
)
|
||||
def test_timestamp_columns_survive_2038_on_postgresql(
|
||||
table_name: str, column_name: str
|
||||
) -> None:
|
||||
"""Unix seconds in INT4 break on 2038-01-19, and long-dated expiries sooner."""
|
||||
column = _table(table_name).columns[column_name]
|
||||
assert _render(column.type, POSTGRES_DIALECT) == "BIGINT"
|
||||
|
||||
|
||||
def test_identifier_columns_stay_narrow() -> None:
|
||||
"""Widening is deliberate, not a blanket sweep over every integer."""
|
||||
for table_name, column_name in [
|
||||
("upstream_providers", "id"),
|
||||
("model_paths", "id"),
|
||||
("model_paths", "upstream_provider_id"),
|
||||
("models", "upstream_provider_id"),
|
||||
("models", "context_length"),
|
||||
]:
|
||||
column = _table(table_name).columns[column_name]
|
||||
assert isinstance(column.type, Integer)
|
||||
assert not isinstance(column.type, BigInteger), (
|
||||
f"{table_name}.{column_name} was widened unnecessarily"
|
||||
)
|
||||
|
||||
|
||||
def test_nsec_state_does_not_require_a_native_postgres_enum() -> None:
|
||||
"""A bare Python Enum makes PostgreSQL demand a ``nsecstate`` TYPE.
|
||||
|
||||
The migration created this column as VARCHAR and never created that type,
|
||||
so every read and write of the secrets singleton — and with it node
|
||||
bootstrap — failed with `type "nsecstate" does not exist`.
|
||||
"""
|
||||
ddl = _render(CreateTable(_table("secrets")), POSTGRES_DIALECT)
|
||||
assert "nsecstate" not in ddl.lower()
|
||||
assert "VARCHAR" in ddl
|
||||
|
||||
column = _table("secrets").columns["nsec_state"]
|
||||
assert column.type.native_enum is False
|
||||
# Values, not member names, so rows written before the column was typed
|
||||
# still load.
|
||||
assert set(column.type.enums) == {"legacy", "encrypted", "cleared"}
|
||||
|
||||
|
||||
def test_every_table_compiles_as_postgresql_ddl() -> None:
|
||||
"""No model may carry DDL that only SQLite can render."""
|
||||
for table in SQLModel.metadata.tables.values():
|
||||
_render(CreateTable(table), POSTGRES_DIALECT)
|
||||
|
||||
|
||||
def test_model_path_upsert_is_built_for_the_bound_dialect() -> None:
|
||||
"""``ON CONFLICT`` is spelled per dialect and the constructs are not swappable.
|
||||
|
||||
The SQLite construct does not compile against PostgreSQL: it raises
|
||||
``AttributeError: 'OnConflictDoUpdate' object has no attribute
|
||||
'constraint_target'``, which killed every model-path refresh.
|
||||
"""
|
||||
from routstr.upstream import model_paths
|
||||
|
||||
class _Bound:
|
||||
def __init__(self, name: str) -> None:
|
||||
self.dialect = type("_D", (), {"name": name})()
|
||||
|
||||
class _Session:
|
||||
def __init__(self, name: str) -> None:
|
||||
self._bind = _Bound(name)
|
||||
|
||||
def get_bind(self) -> Any:
|
||||
return self._bind
|
||||
|
||||
postgres_insert = model_paths._upsert(_Session("postgresql")) # type: ignore[arg-type]
|
||||
sqlite_insert = model_paths._upsert(_Session("sqlite")) # type: ignore[arg-type]
|
||||
|
||||
assert postgres_insert.__class__.__module__.endswith("postgresql.dml")
|
||||
assert sqlite_insert.__class__.__module__.endswith("sqlite.dml")
|
||||
|
||||
# Both must actually compile the upsert they will be asked to run.
|
||||
conflict = dict(
|
||||
index_elements=["model_id", "path", "upstream_provider_id"],
|
||||
set_={"updated_at": postgres_insert.excluded.updated_at},
|
||||
)
|
||||
_render(postgres_insert.on_conflict_do_update(**conflict), POSTGRES_DIALECT)
|
||||
_render(
|
||||
sqlite_insert.on_conflict_do_update(
|
||||
index_elements=conflict["index_elements"],
|
||||
set_={"updated_at": sqlite_insert.excluded.updated_at},
|
||||
),
|
||||
SQLITE_DIALECT,
|
||||
)
|
||||
|
||||
|
||||
def test_clear_alembic_version_uses_the_configured_async_driver() -> None:
|
||||
"""The recovery path must not need a second, sync-only driver.
|
||||
|
||||
It used to build a sync engine from ``DATABASE_URL.replace("+aiosqlite", "")``,
|
||||
which left ``postgresql+asyncpg`` intact and raised ``MissingGreenlet``.
|
||||
"""
|
||||
import inspect
|
||||
|
||||
source = inspect.getsource(db_module._clear_alembic_version)
|
||||
assert "create_engine(" not in source.replace("create_async_engine(", "")
|
||||
assert 'replace("+aiosqlite"' not in source
|
||||
|
||||
|
||||
def test_sqlite_only_maintenance_stays_scoped_to_the_wallet_directory() -> None:
|
||||
"""PRAGMA/sqlite3 handling in db.py may only touch ``.wallet/*.sqlite3``."""
|
||||
import inspect
|
||||
|
||||
source = inspect.getsource(db_module.fix_cashu_migrations)
|
||||
assert ".wallet" in source or "wallet_dir" in source
|
||||
assert "*.sqlite3" in source
|
||||
|
||||
init_source = inspect.getsource(db_module.init_db)
|
||||
# Gated on the bound dialect, not a DATABASE_URL string prefix.
|
||||
assert 'dialect.name == "sqlite"' in init_source
|
||||
@@ -2711,6 +2711,7 @@ source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "aiosqlite" },
|
||||
{ name = "alembic" },
|
||||
{ name = "asyncpg" },
|
||||
{ name = "cashu" },
|
||||
{ name = "fastapi", extra = ["standard-no-fastapi-cloud-cli"] },
|
||||
{ name = "greenlet" },
|
||||
@@ -2747,6 +2748,7 @@ dev = [
|
||||
requires-dist = [
|
||||
{ name = "aiosqlite", specifier = ">=0.20" },
|
||||
{ name = "alembic", specifier = ">=1.13" },
|
||||
{ name = "asyncpg", specifier = ">=0.30" },
|
||||
{ name = "cashu", specifier = ">=0.20" },
|
||||
{ name = "fastapi", extras = ["standard-no-fastapi-cloud-cli"], specifier = ">=0.141" },
|
||||
{ name = "greenlet", specifier = ">=3.2.1" },
|
||||
|
||||
Reference in New Issue
Block a user