mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
auto run migrations on startup
This commit is contained in:
+16
-13
@@ -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())
|
||||
|
||||
@@ -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"}
|
||||
|
||||
@@ -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
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user