import asyncio import logging import os import subprocess from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.pool import NullPool from src.config import Settings, build_database_url logger = logging.getLogger(__name__) settings = Settings() database_url = settings.database_url # SQLite requires aiosqlite and different connect args connect_args = {} if database_url.startswith("sqlite"): connect_args = {"check_same_thread": False} engine = create_async_engine( database_url, future=True, poolclass=NullPool, connect_args=connect_args, ) SessionLocal = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) async def init_database( max_retries: int = 5, retry_delay: float = 2.0, ) -> bool: """Initialize the database by running pending migrations. Uses subprocess to run 'alembic upgrade head' to avoid async/sync context manager issues with SQLAlchemy 2.0. Returns True if migrations succeeded, False otherwise. """ for attempt in range(1, max_retries + 1): try: # Test basic connectivity from sqlalchemy import text test_conn = await engine.connect() try: await test_conn.execute(text("SELECT 1")) finally: await test_conn.close() logger.info("Database connection established.") # Run migrations via subprocess logger.info("Running database migrations...") result = await asyncio.get_event_loop().run_in_executor( None, lambda: subprocess.run( ["python3", "-m", "alembic", "upgrade", "head"], capture_output=True, text=True, cwd=os.path.dirname(os.path.dirname(os.path.abspath(__file__))), ), ) if result.returncode == 0: logger.info("Database migrations completed successfully.") logger.debug("Alembic output: %s", result.stdout) return True else: logger.error("Migration failed: %s", result.stderr) if attempt < max_retries: wait = retry_delay * (2 ** (attempt - 1)) logger.info("Retrying in %.1f seconds...", wait) await asyncio.sleep(wait) else: return False except Exception as exc: error_msg = str(exc).lower() if "connection" in error_msg or "could not connect" in error_msg: logger.warning( "Database connection failed (attempt %d/%d): %s", attempt, max_retries, exc, ) elif "authentication" in error_msg or "password" in error_msg: logger.error( "Database authentication failed: %s. " "Check POSTGRES_USER and POSTGRES_PASSWORD environment variables.", exc, ) return False else: logger.error( "Database initialization error (attempt %d/%d): %s", attempt, max_retries, exc, ) if attempt < max_retries: wait = retry_delay * (2 ** (attempt - 1)) logger.info("Retrying in %.1f seconds...", wait) await asyncio.sleep(wait) else: logger.error( "Failed to initialize database after %d attempts. " "Ensure the database is running and accessible.", max_retries, ) return False return False __all__ = ["SessionLocal", "build_database_url", "engine", "settings", "init_database"]