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:
9qeklajc
2026-09-30 13:20:02 +00:00
parent 8a06e6ab36
commit 250864b45d
8 changed files with 1441 additions and 40 deletions
@@ -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())
+1
View File
@@ -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
View File
@@ -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.",
+18 -2
View File
@@ -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"],
+29
View File
@@ -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 [])
+251
View File
@@ -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
Generated
+2
View File
@@ -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" },