auto run migrations on startup

This commit is contained in:
Shroominic
2025-08-09 13:50:05 -03:00
parent 12d2d3714f
commit 2698d0fa81
4 changed files with 75 additions and 18 deletions
+16 -13
View File
@@ -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())
+4 -4
View File
@@ -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"}
+46
View File
@@ -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
+9 -1
View File
@@ -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())