Files
live-hub-py/web/backend/database.py
T

143 lines
4.3 KiB
Python

"""数据库引擎与会话管理"""
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]
DB_PATH = PROJECT_ROOT / "data" / "web.db"
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
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"
)
return f"sqlite:///{DB_PATH}"
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)
command.upgrade(config, "head")
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} 个值")