seed_dev.py 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190
  1. """显式、幂等的本地最小开发数据初始化。
  2. 执行:python -m dms.seed_dev
  3. 必须通过环境变量提供三个用户密码;默认不会覆盖已有用户密码。
  4. """
  5. from __future__ import annotations
  6. import os
  7. from flask import Flask
  8. from sqlalchemy import select
  9. from dms import init_dms
  10. from dms.common.enums import EnabledStatus, RoleCode, SecurityLevel
  11. from dms.extensions import db
  12. from dms.models import Organization, User
  13. from dms.security.passwords import hash_password
  14. PASSWORD_ENVIRONMENTS = {
  15. "admin": "DMS_SEED_ADMIN_PASSWORD",
  16. "auditor": "DMS_SEED_AUDITOR_PASSWORD",
  17. "user": "DMS_SEED_USER_PASSWORD",
  18. }
  19. def _required_passwords() -> dict[str, str]:
  20. missing = [
  21. environment
  22. for environment in PASSWORD_ENVIRONMENTS.values()
  23. if not os.environ.get(environment)
  24. ]
  25. if missing:
  26. raise RuntimeError("缺少必需的密码环境变量:" + "、".join(missing))
  27. return {
  28. username: os.environ[environment]
  29. for username, environment in PASSWORD_ENVIRONMENTS.items()
  30. }
  31. def _overwrite_passwords_enabled() -> bool:
  32. raw_value = os.environ.get("DMS_SEED_OVERWRITE_PASSWORDS", "false").lower()
  33. if raw_value not in {"true", "false"}:
  34. raise RuntimeError("DMS_SEED_OVERWRITE_PASSWORDS必须是true或false")
  35. return raw_value == "true"
  36. def _organization(
  37. *,
  38. org_code: str,
  39. org_name: str,
  40. org_path: str,
  41. sort_no: int,
  42. parent_id: int | None,
  43. ) -> tuple[Organization, bool]:
  44. existing = db.session.scalar(
  45. select(Organization).where(Organization.org_code == org_code)
  46. )
  47. if existing is not None:
  48. return existing, False
  49. organization = Organization(
  50. org_code=org_code,
  51. org_name=org_name,
  52. org_path=org_path,
  53. sort_no=sort_no,
  54. parent_id=parent_id,
  55. status=EnabledStatus.ENABLED.value,
  56. )
  57. db.session.add(organization)
  58. db.session.flush()
  59. return organization, True
  60. def _user(
  61. *,
  62. username: str,
  63. real_name: str,
  64. password: str,
  65. role_code: RoleCode,
  66. security_level: SecurityLevel,
  67. organization: Organization,
  68. overwrite_password: bool,
  69. ) -> bool:
  70. existing = db.session.scalar(select(User).where(User.username == username))
  71. if existing is not None:
  72. if overwrite_password:
  73. existing.password_hash = hash_password(password)
  74. existing.auth_version += 1
  75. existing.row_version += 1
  76. return False
  77. db.session.add(
  78. User(
  79. username=username,
  80. password_hash=hash_password(password),
  81. real_name=real_name,
  82. organization_id=organization.id,
  83. organization_name=organization.org_name,
  84. role_code=role_code.value,
  85. security_level=security_level.value,
  86. status=EnabledStatus.ENABLED.value,
  87. )
  88. )
  89. return True
  90. def seed() -> tuple[int, int]:
  91. passwords = _required_passwords()
  92. overwrite_passwords = _overwrite_passwords_enabled()
  93. created_organizations = 0
  94. created_users = 0
  95. root, created = _organization(
  96. org_code="ORG_ROOT",
  97. org_name="机关",
  98. org_path="/机关",
  99. sort_no=10,
  100. parent_id=None,
  101. )
  102. created_organizations += int(created)
  103. ops, created = _organization(
  104. org_code="ORG_OPS",
  105. org_name="作战部",
  106. org_path="/机关/作战部",
  107. sort_no=20,
  108. parent_id=root.id,
  109. )
  110. created_organizations += int(created)
  111. _, created = _organization(
  112. org_code="ORG_COMMS",
  113. org_name="通信部",
  114. org_path="/机关/通信部",
  115. sort_no=30,
  116. parent_id=root.id,
  117. )
  118. created_organizations += int(created)
  119. created_users += int(
  120. _user(
  121. username="admin",
  122. real_name="系统管理员",
  123. password=passwords["admin"],
  124. role_code=RoleCode.ADMIN,
  125. security_level=SecurityLevel.TOP_SECRET,
  126. organization=root,
  127. overwrite_password=overwrite_passwords,
  128. )
  129. )
  130. created_users += int(
  131. _user(
  132. username="auditor",
  133. real_name="审计员",
  134. password=passwords["auditor"],
  135. role_code=RoleCode.AUDITOR,
  136. security_level=SecurityLevel.TOP_SECRET,
  137. organization=root,
  138. overwrite_password=overwrite_passwords,
  139. )
  140. )
  141. created_users += int(
  142. _user(
  143. username="user",
  144. real_name="普通用户",
  145. password=passwords["user"],
  146. role_code=RoleCode.USER,
  147. security_level=SecurityLevel.SECRET,
  148. organization=ops,
  149. overwrite_password=overwrite_passwords,
  150. )
  151. )
  152. db.session.commit()
  153. return created_organizations, created_users
  154. def main() -> None:
  155. app = Flask("dms-seed-dev")
  156. init_dms(app)
  157. if _overwrite_passwords_enabled():
  158. print("警告:已显式启用已有开发用户密码覆盖,并将使其旧Token失效。")
  159. with app.app_context():
  160. try:
  161. organization_count, user_count = seed()
  162. except Exception:
  163. db.session.rollback()
  164. raise
  165. print(f"初始化完成:新增组织{organization_count}个,新增用户{user_count}个。")
  166. if __name__ == "__main__":
  167. main()