"""数据库自动初始化模块。 在应用启动时自动执行: 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) # 种子数据失败不阻止启动