From 2698d0fa8122bbb3eb892baf0e2db8139c9b1c70 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Sat, 9 Aug 2025 13:50:05 -0300 Subject: [PATCH] auto run migrations on startup --- migrations/env.py | 29 +++++++++++++----------- migrations/script.py.mako | 8 +++---- router/core/db.py | 46 +++++++++++++++++++++++++++++++++++++++ router/core/main.py | 10 ++++++++- 4 files changed, 75 insertions(+), 18 deletions(-) diff --git a/migrations/env.py b/migrations/env.py index 0134da9b..e62499fa 100644 --- a/migrations/env.py +++ b/migrations/env.py @@ -1,6 +1,6 @@ import asyncio -import importlib.util import pathlib +import sys from logging.config import fileConfig from alembic import context @@ -9,18 +9,10 @@ from sqlalchemy.engine import Connection from sqlalchemy.ext.asyncio import create_async_engine from sqlmodel import SQLModel -db_path = pathlib.Path(__file__).resolve().parents[1] / "router" / "core" / "db.py" -spec = importlib.util.spec_from_file_location("core.db", db_path) -if spec is None: - raise ImportError(f"Could not load spec from {db_path}") -db = importlib.util.module_from_spec(spec) -if db is None: - raise ImportError(f"Could not load module from {db_path}") -if spec.loader is None: - raise ImportError(f"Spec loader is None for {db_path}") -spec.loader.exec_module(db) +# Add the parent directory to the Python path so we can import router modules +sys.path.insert(0, str(pathlib.Path(__file__).resolve().parents[1])) -DATABASE_URL = getattr(db, "DATABASE_URL", "sqlite+aiosqlite:///keys.db") +from router.core.db import DATABASE_URL config = context.config if config.config_file_name is None: @@ -64,4 +56,15 @@ async def run_migrations_online() -> None: if context.is_offline_mode(): run_migrations_offline() else: - asyncio.run(run_migrations_online()) + # Check if we're already in an event loop (e.g., being called from FastAPI) + try: + loop = asyncio.get_running_loop() + # If we're in an existing loop, create a new thread to run migrations + import concurrent.futures + + with concurrent.futures.ThreadPoolExecutor() as executor: + future = executor.submit(asyncio.run, run_migrations_online()) + future.result() + except RuntimeError: + # No event loop running, we can use asyncio.run directly + asyncio.run(run_migrations_online()) diff --git a/migrations/script.py.mako b/migrations/script.py.mako index 12efbe18..50ab7623 100644 --- a/migrations/script.py.mako +++ b/migrations/script.py.mako @@ -5,20 +5,20 @@ Revision ID: ${up_revision} Revises: ${down_revision | comma,n} Create Date: ${create_date} """ + import sqlalchemy as sa -import sqlmodel as sqlm +import sqlmodel from alembic import op ${imports if imports else ""} - # revision identifiers, used by Alembic. revision = ${repr(up_revision)} down_revision = ${repr(down_revision)} branch_labels = ${repr(branch_labels)} depends_on = ${repr(depends_on)} -def upgrade(): +def upgrade() -> None: ${upgrades if upgrades else "pass"} -def downgrade(): +def downgrade() -> None: ${downgrades if downgrades else "pass"} diff --git a/router/core/db.py b/router/core/db.py index e959e6b4..ac94b291 100644 --- a/router/core/db.py +++ b/router/core/db.py @@ -2,10 +2,16 @@ import os from contextlib import asynccontextmanager from typing import AsyncGenerator +from alembic import command +from alembic.config import Config from sqlalchemy.ext.asyncio.engine import create_async_engine from sqlmodel import Field, SQLModel from sqlmodel.ext.asyncio.session import AsyncSession +from .logging import get_logger + +logger = get_logger(__name__) + DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") @@ -29,6 +35,10 @@ class ApiKey(SQLModel, table=True): # type: ignore default=0, description="Total spent in millisatoshis (msats)" ) total_requests: int = Field(default=0) + mint_url: str | None = Field( + default=None, + description="URL of the mint used to create the cashu-token", + ) async def init_db() -> None: @@ -46,3 +56,39 @@ async def get_session() -> AsyncGenerator[AsyncSession, None]: async def create_session() -> AsyncGenerator[AsyncSession, None]: async with AsyncSession(engine, expire_on_commit=False) as session: yield session + + +def run_migrations() -> None: + """Run Alembic migrations programmatically.""" + import pathlib + + try: + logger.info("Starting database migrations") + + # Get the path to the alembic.ini file + project_root = pathlib.Path(__file__).resolve().parents[2] + alembic_ini_path = project_root / "alembic.ini" + + if not alembic_ini_path.exists(): + raise FileNotFoundError( + f"Alembic configuration file not found at {alembic_ini_path}" + ) + + # Create Alembic config object + alembic_cfg = Config(str(alembic_ini_path)) + + # Set the database URL in the config + alembic_cfg.set_main_option("sqlalchemy.url", DATABASE_URL) + + # Run migrations to the latest revision + logger.info("Running migrations to latest revision") + command.upgrade(alembic_cfg, "head") + + logger.info("Database migrations completed successfully") + + except Exception as e: + logger.error( + "Database migration failed", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + raise diff --git a/router/core/main.py b/router/core/main.py index e7b77e91..2e0dce68 100644 --- a/router/core/main.py +++ b/router/core/main.py @@ -12,7 +12,7 @@ from ..payment.models import MODELS, models_router, update_sats_pricing from ..proxy import proxy_router from ..wallet import periodic_payout from .admin import admin_router -from .db import init_db +from .db import init_db, run_migrations from .logging import get_logger, setup_logging # Initialize logging first @@ -30,6 +30,14 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: payout_task = None try: + # Run database migrations on startup + # This ensures the database schema is always up-to-date in production + # Migrations are idempotent - running them multiple times is safe + logger.info("Running database migrations") + run_migrations() + + # Initialize database connection pools + # This creates any tables that might not be tracked by migrations yet await init_db() pricing_task = asyncio.create_task(update_sats_pricing())