seed_dev.py 4.9 KB

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