diff --git a/migrations/versions/d3c8b21f7a04_widen_money_and_timestamp_columns.py b/migrations/versions/d3c8b21f7a04_widen_money_and_timestamp_columns.py new file mode 100644 index 00000000..8022988f --- /dev/null +++ b/migrations/versions/d3c8b21f7a04_widen_money_and_timestamp_columns.py @@ -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()) diff --git a/pyproject.toml b/pyproject.toml index 296fab3f..c26bf316 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/routstr/core/db.py b/routstr/core/db.py index 500420d5..725ef25f 100644 --- a/routstr/core/db.py +++ b/routstr/core/db.py @@ -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.", diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index a53ec0f0..6b212f52 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -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"], diff --git a/tests/integration/README.md b/tests/integration/README.md index a25e1b54..bfe3f072 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -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 diff --git a/tests/integration/test_postgres_compatibility.py b/tests/integration/test_postgres_compatibility.py new file mode 100644 index 00000000..d9f6b5be --- /dev/null +++ b/tests/integration/test_postgres_compatibility.py @@ -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 []) diff --git a/tests/unit/test_postgres_schema_compat.py b/tests/unit/test_postgres_schema_compat.py new file mode 100644 index 00000000..f0229f1a --- /dev/null +++ b/tests/unit/test_postgres_schema_compat.py @@ -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 diff --git a/uv.lock b/uv.lock index 86d2be5d..f7ea8c6e 100644 --- a/uv.lock +++ b/uv.lock @@ -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" },