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 asyncio
|
||||||
import importlib.util
|
|
||||||
import pathlib
|
import pathlib
|
||||||
|
import sys
|
||||||
from logging.config import fileConfig
|
from logging.config import fileConfig
|
||||||
|
|
||||||
from alembic import context
|
from alembic import context
|
||||||
@@ -9,18 +9,10 @@ from sqlalchemy.engine import Connection
|
|||||||
from sqlalchemy.ext.asyncio import create_async_engine
|
from sqlalchemy.ext.asyncio import create_async_engine
|
||||||
from sqlmodel import SQLModel
|
from sqlmodel import SQLModel
|
||||||
|
|
||||||
db_path = pathlib.Path(__file__).resolve().parents[1] / "router" / "core" / "db.py"
|
# Add the parent directory to the Python path so we can import router modules
|
||||||
spec = importlib.util.spec_from_file_location("core.db", db_path)
|
sys.path.insert(0, str(pathlib.Path(__file__).resolve().parents[1]))
|
||||||
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)
|
|
||||||
|
|
||||||
DATABASE_URL = getattr(db, "DATABASE_URL", "sqlite+aiosqlite:///keys.db")
|
from router.core.db import DATABASE_URL
|
||||||
|
|
||||||
config = context.config
|
config = context.config
|
||||||
if config.config_file_name is None:
|
if config.config_file_name is None:
|
||||||
@@ -64,4 +56,15 @@ async def run_migrations_online() -> None:
|
|||||||
if context.is_offline_mode():
|
if context.is_offline_mode():
|
||||||
run_migrations_offline()
|
run_migrations_offline()
|
||||||
else:
|
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}
|
Revises: ${down_revision | comma,n}
|
||||||
Create Date: ${create_date}
|
Create Date: ${create_date}
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import sqlalchemy as sa
|
import sqlalchemy as sa
|
||||||
import sqlmodel as sqlm
|
import sqlmodel
|
||||||
from alembic import op
|
from alembic import op
|
||||||
${imports if imports else ""}
|
${imports if imports else ""}
|
||||||
|
|
||||||
# revision identifiers, used by Alembic.
|
# revision identifiers, used by Alembic.
|
||||||
revision = ${repr(up_revision)}
|
revision = ${repr(up_revision)}
|
||||||
down_revision = ${repr(down_revision)}
|
down_revision = ${repr(down_revision)}
|
||||||
branch_labels = ${repr(branch_labels)}
|
branch_labels = ${repr(branch_labels)}
|
||||||
depends_on = ${repr(depends_on)}
|
depends_on = ${repr(depends_on)}
|
||||||
|
|
||||||
def upgrade():
|
def upgrade() -> None:
|
||||||
${upgrades if upgrades else "pass"}
|
${upgrades if upgrades else "pass"}
|
||||||
|
|
||||||
|
|
||||||
def downgrade():
|
def downgrade() -> None:
|
||||||
${downgrades if downgrades else "pass"}
|
${downgrades if downgrades else "pass"}
|
||||||
|
|||||||
@@ -2,10 +2,16 @@ import os
|
|||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
|
from alembic import command
|
||||||
|
from alembic.config import Config
|
||||||
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
||||||
from sqlmodel import Field, SQLModel
|
from sqlmodel import Field, SQLModel
|
||||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
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")
|
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)"
|
default=0, description="Total spent in millisatoshis (msats)"
|
||||||
)
|
)
|
||||||
total_requests: int = Field(default=0)
|
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:
|
async def init_db() -> None:
|
||||||
@@ -46,3 +56,39 @@ async def get_session() -> AsyncGenerator[AsyncSession, None]:
|
|||||||
async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||||
async with AsyncSession(engine, expire_on_commit=False) as session:
|
async with AsyncSession(engine, expire_on_commit=False) as session:
|
||||||
yield 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 ..proxy import proxy_router
|
||||||
from ..wallet import periodic_payout
|
from ..wallet import periodic_payout
|
||||||
from .admin import admin_router
|
from .admin import admin_router
|
||||||
from .db import init_db
|
from .db import init_db, run_migrations
|
||||||
from .logging import get_logger, setup_logging
|
from .logging import get_logger, setup_logging
|
||||||
|
|
||||||
# Initialize logging first
|
# Initialize logging first
|
||||||
@@ -30,6 +30,14 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
|||||||
payout_task = None
|
payout_task = None
|
||||||
|
|
||||||
try:
|
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()
|
await init_db()
|
||||||
|
|
||||||
pricing_task = asyncio.create_task(update_sats_pricing())
|
pricing_task = asyncio.create_task(update_sats_pricing())
|
||||||
|
|||||||
Reference in New Issue
Block a user