144 lines
4.3 KiB
Python
144 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]
|
||
|
||
|
||
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)
|
||
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} 个值")
|