dingdanquanliucheng/backend/app/startup/db_init.py

266 lines
8.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""数据库自动初始化模块。
在应用启动时自动执行:
1. 检查并创建数据库(如果不存在)
2. 执行 Alembic 迁移,确保表结构最新
3. 补全 ORM 模型中定义但数据库不存在的表
4. 补全已有表中缺失的字段(只 ADD COLUMN不 MODIFY/DELETE
5. 执行种子数据初始化
被调用方backend.app.main 的 lifespan 函数。
"""
import logging
import subprocess
import sys
from pathlib import Path
from sqlalchemy import create_engine, text
from sqlalchemy.exc import OperationalError
from backend.app.core.config import get_settings
logger = logging.getLogger(__name__)
PROJECT_ROOT = Path(__file__).resolve().parents[3]
def auto_init_database() -> None:
"""自动初始化数据库。
启动时调用,按顺序执行:
1. 检查数据库是否存在,不存在则创建
2. 执行 Alembic 迁移
3. 检查并补全所有表Alembic 迁移可能遗漏新增模型)
4. 执行种子数据初始化
"""
settings = get_settings()
# 1. 确保数据库存在
_ensure_database_exists(settings)
# 2. 执行 Alembic 迁移
_run_alembic_migrations()
# 3. 补全检查:确保所有 ORM 模型对应的表都已创建
_create_tables_directly()
# 4. 补全已有表的缺失字段
_sync_columns()
# 5. 初始化种子数据
_run_bootstrap_data()
def _ensure_database_exists(settings) -> None:
"""检查数据库是否存在,不存在则创建。"""
# 先尝试连接目标数据库
db_url = (
f"mysql+pymysql://{settings.mysql_user}:{settings.mysql_password}"
f"@{settings.mysql_host}:{settings.mysql_port}/{settings.mysql_database}"
)
engine = create_engine(db_url)
try:
with engine.connect() as conn:
conn.execute(text("SELECT 1"))
logger.info("数据库 '%s' 已存在,连接正常", settings.mysql_database)
engine.dispose()
return
except OperationalError as e:
# 如果是"未知数据库"错误,则创建数据库
if "Unknown database" in str(e) or "1049" in str(e):
logger.info("数据库 '%s' 不存在,正在创建...", settings.mysql_database)
engine.dispose()
_create_database(settings)
else:
logger.error("数据库连接失败: %s", e)
raise
def _create_database(settings) -> None:
"""创建数据库。"""
# 连接到 MySQL 服务器(不指定数据库)
server_url = (
f"mysql+pymysql://{settings.mysql_user}:{settings.mysql_password}"
f"@{settings.mysql_host}:{settings.mysql_port}"
)
engine = create_engine(server_url)
try:
with engine.connect() as conn:
conn.execute(text(
f"CREATE DATABASE IF NOT EXISTS `{settings.mysql_database}` "
f"CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"
))
conn.commit()
logger.info("数据库 '%s' 创建成功", settings.mysql_database)
except Exception as e:
logger.error("创建数据库失败: %s", e)
raise
finally:
engine.dispose()
def _run_alembic_migrations() -> None:
"""执行 Alembic 迁移,确保表结构最新。"""
logger.info("正在执行数据库迁移...")
alembic_ini = PROJECT_ROOT / "backend" / "alembic.ini"
if not alembic_ini.exists():
logger.warning("未找到 alembic.ini尝试使用 SQLAlchemy 直接创建表")
_create_tables_directly()
return
try:
result = subprocess.run(
[sys.executable, "-m", "alembic", "upgrade", "head"],
cwd=str(PROJECT_ROOT / "backend"),
capture_output=True,
text=True,
timeout=60,
)
if result.returncode == 0:
logger.info("数据库迁移完成")
if result.stdout.strip():
logger.debug("迁移输出: %s", result.stdout.strip())
else:
logger.error("数据库迁移失败,尝试直接创建表: %s", result.stderr[:200])
_create_tables_directly()
except subprocess.TimeoutExpired:
logger.error("数据库迁移超时,尝试直接创建表")
_create_tables_directly()
except Exception as e:
logger.error("执行迁移异常,尝试直接创建表: %s", e)
_create_tables_directly()
def _create_tables_directly() -> None:
"""检查并补全所有 ORM 模型对应的数据库表。
Base.metadata.create_all 默认 checkfirst=True
只创建不存在的表,已存在的表不会被修改或删除。
"""
try:
from sqlalchemy import inspect
from backend.app.db import Base, engine
# 显式导入所有模型类,确保每个 ORM 类都注册到 Base.metadata
from backend.app.models import ( # noqa: F401
AIRecognitionLog,
AuditLog,
Customer,
CustomerArrears,
FileAttachment,
LogisticsTask,
LogisticsTrace,
LogisticsWaybill,
Menu,
OrderSupplierTextLog,
PerformanceStatCache,
Product,
ProductCategory,
ProductPriceTier,
ProductPricingRule,
Role,
RoleMenu,
SalesOrder,
SalesOrderApproveLog,
SalesOrderItem,
Supplier,
SupplierProductCost,
SystemConfig,
SystemReminder,
User,
)
orm_tables = set(Base.metadata.tables.keys())
logger.info("ORM 已注册 %d 张表: %s", len(orm_tables), sorted(orm_tables))
# 记录创建前已存在的表
inspector = inspect(engine)
before = set(inspector.get_table_names())
# create_all: 已存在的跳过,不存在的创建
Base.metadata.create_all(engine, checkfirst=True)
# 对比前后差异,确认结果
after = set(inspect(engine).get_table_names())
newly_created = (orm_tables - before) & after
still_missing = orm_tables - after
if newly_created:
logger.info("新创建 %d 张表: %s", len(newly_created), sorted(newly_created))
if still_missing:
logger.error("以下表仍然缺失,请检查: %s", sorted(still_missing))
else:
logger.info("所有 %d 张表均已就绪", len(orm_tables))
except Exception as e:
logger.error("检查/创建表失败: %s", e)
def _sync_columns() -> None:
"""检查已有表是否缺少 ORM 模型中定义的字段,自动补全。
只做 ADD COLUMN不做 MODIFY/DELETE确保安全。
"""
try:
from sqlalchemy import inspect, text
from backend.app.db import Base, engine
inspector = inspect(engine)
existing_tables = set(inspector.get_table_names())
added_count = 0
for table_name, table in Base.metadata.tables.items():
if table_name not in existing_tables:
continue
db_columns = {col["name"] for col in inspector.get_columns(table_name)}
for col in table.columns:
if col.name not in db_columns:
# 生成 ALTER TABLE ADD COLUMN 语句
col_type = col.type.compile(dialect=engine.dialect)
nullable = "NULL" if col.nullable else "NOT NULL"
default = ""
if col.default is not None:
default_val = col.default.arg
if callable(default_val):
# 函数默认值(如 now())跳过,让数据库处理
default = ""
elif isinstance(default_val, str):
default = f"DEFAULT '{default_val}'"
else:
default = f"DEFAULT {default_val}"
sql = f"ALTER TABLE `{table_name}` ADD COLUMN `{col.name}` {col_type} {nullable} {default}"
logger.info("补字段: %s.%s", table_name, col.name)
with engine.connect() as conn:
conn.execute(text(sql))
conn.commit()
added_count += 1
if added_count:
logger.info("共补全 %d 个缺失字段", added_count)
else:
logger.info("所有表字段均已同步")
except Exception as e:
logger.error("同步表字段失败: %s", e)
def _run_bootstrap_data() -> None:
"""执行种子数据初始化。"""
logger.info("正在初始化种子数据...")
try:
from backend.scripts.bootstrap_data import main as bootstrap_main
bootstrap_main()
logger.info("种子数据初始化完成")
except Exception as e:
logger.error("种子数据初始化失败: %s", e)
# 种子数据失败不阻止启动