Files
dpb/backend/app/scripts/initialize.py
T
34047007@qq.com b95053c52c init: 初始化 dpb 桃育种系统代码库
前后端 + 后端 FastAPI 全量源码、部署脚本与文档。
2026-08-06 00:17:49 +08:00

170 lines
6.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.
import json
import secrets
from typing import Any
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.v1.module_system.dept.model import DeptModel
from app.api.v1.module_system.dict.model import DictDataModel, DictTypeModel
from app.api.v1.module_system.menu.model import MenuModel
from app.api.v1.module_system.params.model import ParamsModel
from app.api.v1.module_system.role.model import RoleModel
from app.api.v1.module_bre.germplasm.model import BreedingGermplasmModel
from app.api.v1.module_system.user.model import UserModel, UserRolesModel
from app.api.v1.module_system.versions.model import VersionModel
from app.common.enums import EnvironmentEnum
from app.config.path_conf import LOG_DIR, SCRIPT_DIR
from app.config.setting import settings
from app.core.database import async_db_session, check_db, create_tables
from app.core.logger import logger
from app.utils.password_util import PwdUtil
class InitializeData:
"""初始化数据库和基础数据"""
# 按依赖关系排序:先基础表,再关联表
prepare_init_models: list[type] = [
MenuModel,
DeptModel,
ParamsModel,
RoleModel,
DictTypeModel,
DictDataModel,
UserModel,
UserRolesModel,
VersionModel,
BreedingGermplasmModel,
]
# 树形模型:JSON 含嵌套 children,需递归创建对象
_RECURSIVE_TABLES: set[str] = {"sys_menu", "sys_dept"}
async def init_db(self) -> None:
"""建表并导入种子数据"""
await check_db()
# await drop_tables()
await create_tables()
async with async_db_session() as session, session.begin():
await self.__init_data(session)
async def __init_data(self, db: AsyncSession) -> None:
"""按依赖顺序初始化各表种子数据"""
dict_type_mapping: dict[str, Any] = {}
for model in self.prepare_init_models:
table_name = model.__tablename__
data = await self.__load_json(table_name)
if not data:
logger.info(f"⏭️ 跳过 {table_name} 表,无初始化数据")
continue
# 已有数据则跳过
count = await db.execute(select(func.count()).select_from(model))
if count.scalar():
logger.info(f"⏭️ 跳过 {table_name} 表数据初始化(表已有数据)")
continue
try:
if table_name in self._RECURSIVE_TABLES:
objs = self.__create_objects_with_children(data, model)
elif table_name == "sys_dict_type":
objs = []
for item in data:
obj = model(**item)
objs.append(obj)
dict_type_mapping[item["dict_type"]] = obj
elif table_name == "sys_dict_data":
objs = []
for item in data:
dict_type_str = item.get("dict_type")
if dict_type_str not in dict_type_mapping:
logger.warning(f"⚠️ 未找到字典类型 {dict_type_str},跳过")
continue
item["dict_type_id"] = dict_type_mapping[dict_type_str].id
objs.append(model(**item))
else:
# 生产首次初始化:把仓库内置的固定种子密码(如 123456)随机化,防止已知密码登录
if settings.ENVIRONMENT == EnvironmentEnum.PROD and table_name == "sys_user":
self._randomize_seed_passwords(data)
objs = [model(**item) for item in data]
if objs:
db.add_all(objs)
await db.flush()
logger.info(f"✅️ 已向 {table_name} 写入初始化数据")
else:
logger.info(f"⏭️ 跳过 {table_name} 表数据初始化(无有效数据)")
except Exception:
logger.error(f"❌️ 初始化 {table_name} 表数据失败")
raise
def _randomize_seed_passwords(self, data: list[dict]) -> None:
"""生产环境首次初始化时,把种子账号的固定弱密码替换为随机值。
仓库内 sys_user.json 携带对所有部署相同的初始密码(如 123456),
直接使用会让任何未改密的实例可被已知密码直接登录。生产初始化时随机化,
一次性初始密码写入日志目录文件,供运维首次登录使用;登录修改后务必删除该文件。
"""
creds: list[str] = []
for item in data:
if not item.get("password"):
continue
raw = secrets.token_urlsafe(18)
item["password"] = PwdUtil.hash_password(raw)
creds.append(f"{item.get('username')}={raw}")
if not creds:
return
try:
LOG_DIR.mkdir(parents=True, exist_ok=True)
note = LOG_DIR / "prod_initial_passwords.txt"
note.write_text(
"生产环境首次初始化种子账号随机初始密码(登录后请修改并删除本文件):\n"
+ "\n".join(creds)
+ "\n",
encoding="utf-8",
)
logger.warning(
"⚠️ 生产初始化已为种子账号生成随机初始密码,已写入 {},请立即登录修改:{}",
note,
"".join(creds),
)
except Exception as e:
logger.error("生成种子初始密码失败: {}", e)
@staticmethod
def __create_objects_with_children(data: list[dict], model_class: type) -> list:
"""递归创建树形模型实例,处理嵌套 children 并注入 parent_id"""
def _create(obj_data: dict) -> Any:
children_data = obj_data.pop("children", [])
obj = model_class(**obj_data)
# 子节点通过 relationship 自动设置 parent_id
if children_data:
obj.children = [_create(child) for child in children_data]
return obj
return [_create(item) for item in data]
async def __load_json(self, filename: str) -> list[dict]:
"""读取并解析种子数据 JSON 文件"""
json_path = SCRIPT_DIR / f"{filename}.json"
if not json_path.exists():
return []
try:
with open(json_path, encoding="utf-8") as f:
return json.load(f)
except json.JSONDecodeError as e:
logger.error(f"❌️ 解析 {json_path} 失败: {e!s}")
raise
except Exception as e:
logger.error(f"❌️ 读取 {json_path} 失败: {e!s}")
raise