"""数据库引擎与会话管理""" import os from pathlib import Path from urllib.parse import quote_plus from sqlalchemy import create_engine, event from sqlalchemy.orm import sessionmaker, declarative_base PROJECT_ROOT = Path(__file__).resolve().parents[2] def _env_int(name: str, default: int) -> int: """读取正整数环境变量,非法值回退到默认值。""" try: value = int(os.getenv(name, str(default))) except ValueError: return default return value if value > 0 else default def _get_database_url() -> str: """优先使用完整连接串,否则按 DB_* 环境变量构造 MySQL 连接。""" database_url = os.getenv("DATABASE_URL", "").strip() if database_url: return database_url db_host = os.getenv("DB_HOST", "").strip() if db_host: db_port = _env_int("DB_PORT", 3306) db_name = os.getenv("DB_NAME", "douyu_login").strip() db_user = os.getenv("DB_USER", "douyu_login").strip() db_password = os.getenv("DB_PASSWORD", "") if not db_name or not db_user or not db_password: raise RuntimeError("使用 DB_HOST 时必须同时设置 DB_NAME、DB_USER 和 DB_PASSWORD") return ( f"mysql+pymysql://{quote_plus(db_user)}:{quote_plus(db_password)}" f"@{db_host}:{db_port}/{db_name}?charset=utf8mb4" ) raise RuntimeError( "未配置数据库连接。请设置 DATABASE_URL=mysql+pymysql://...," "或设置 DB_HOST、DB_NAME、DB_USER、DB_PASSWORD" ) DATABASE_URL = _get_database_url() connect_args = ( {"check_same_thread": False, "timeout": 30} if DATABASE_URL.startswith("sqlite") else {} ) engine_options = { "connect_args": connect_args, "echo": False, } if not DATABASE_URL.startswith("sqlite"): engine_options.update( pool_pre_ping=True, pool_size=_env_int("DB_POOL_SIZE", 20), max_overflow=_env_int("DB_MAX_OVERFLOW", 20), pool_timeout=_env_int("DB_POOL_TIMEOUT", 30), pool_recycle=_env_int("DB_POOL_RECYCLE", 1800), ) engine = create_engine(DATABASE_URL, **engine_options) if DATABASE_URL.startswith("sqlite"): @event.listens_for(engine, "connect") def _set_sqlite_pragmas(dbapi_connection, connection_record): """提升 SQLite 并发写入稳定性。""" cursor = dbapi_connection.cursor() cursor.execute("PRAGMA journal_mode=WAL") cursor.execute("PRAGMA busy_timeout=30000") cursor.execute("PRAGMA foreign_keys=ON") cursor.close() SessionLocal = sessionmaker(bind=engine, autocommit=False, autoflush=False) Base = declarative_base() def get_db(): """FastAPI 依赖:提供数据库会话,请求结束自动关闭。""" db = SessionLocal() try: yield db finally: db.close() def init_db(): """执行数据库迁移 + 写入初始数据。""" run_migrations() _seed() _encrypt_existing_sensitive_data() def run_migrations(): """运行 Alembic 迁移到最新版本。""" from alembic import command from alembic.config import Config config = Config(str(PROJECT_ROOT / "alembic.ini")) config.set_main_option("script_location", str(PROJECT_ROOT / "web" / "backend" / "migrations")) config.set_main_option("sqlalchemy.url", DATABASE_URL) # 标记为应用内嵌调用,env.py 据此跳过 fileConfig,避免覆盖 uvicorn 日志配置。 os.environ["ALEMBIC_EMBEDDED"] = "1" try: command.upgrade(config, "head") finally: os.environ.pop("ALEMBIC_EMBEDDED", None) def _seed(): """写入默认超管账号和角色(账号密码通过环境变量配置)。""" from .models import User from .security import hash_password admin_username = os.getenv("ADMIN_USERNAME", "admin") admin_password = os.getenv("ADMIN_PASSWORD", "admin123") db = SessionLocal() try: if not db.query(User).first(): admin = User( username=admin_username, password_hash=hash_password(admin_password), role="super_admin", is_active=True, remark="默认超级管理员", ) db.add(admin) db.commit() finally: db.close() def _encrypt_existing_sensitive_data(): """启动时把历史明文敏感数据迁移为密文。""" from loguru import logger from .crypto_storage import encrypt_existing_sensitive_data changed = encrypt_existing_sensitive_data(engine) if changed: logger.info(f"已加密历史敏感字段: {changed} 个值")