repositories.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367
  1. from sqlalchemy import select
  2. from sqlalchemy.orm import Session, sessionmaker
  3. from app.core.identifiers import new_ulid
  4. from app.domains.catalog.models import Plan, Product, ProductChangeLog, ProductVersion
  5. from app.domains.catalog.repository import CatalogRepository
  6. from app.domains.identity.models import AdminRole, AdminUser, AuthSession, H5User
  7. from app.domains.identity.repository import IdentityRepository
  8. from app.infrastructure.mysql.core_models import (
  9. AdminUserRecord,
  10. AuthSessionRecord,
  11. H5UserRecord,
  12. PermissionRecord,
  13. PlanRecord,
  14. ProductChangeLogRecord,
  15. ProductRecord,
  16. ProductVersionRecord,
  17. RoleRecord,
  18. )
  19. SessionFactory = sessionmaker[Session]
  20. class SqlAlchemyIdentityRepository(IdentityRepository):
  21. def __init__(self, session_factory: SessionFactory) -> None:
  22. self._session_factory = session_factory
  23. def get_h5_user_by_mobile(self, mobile: str) -> H5User | None:
  24. with self._session_factory() as session:
  25. record = session.scalar(select(H5UserRecord).where(H5UserRecord.mobile == mobile))
  26. return self._to_h5_user(record) if record else None
  27. def get_h5_user(self, user_id: str) -> H5User | None:
  28. with self._session_factory() as session:
  29. record = session.get(H5UserRecord, user_id)
  30. return self._to_h5_user(record) if record else None
  31. def list_h5_users(self) -> list[H5User]:
  32. with self._session_factory() as session:
  33. records = session.scalars(
  34. select(H5UserRecord).order_by(H5UserRecord.created_at.desc())
  35. ).all()
  36. return [self._to_h5_user(record) for record in records]
  37. def save_h5_user(self, user: H5User) -> None:
  38. with self._session_factory() as session:
  39. session.merge(
  40. H5UserRecord(
  41. id=user.id,
  42. mobile=user.mobile,
  43. mobile_masked=user.mobile_masked,
  44. display_name=user.display_name,
  45. status=user.status,
  46. created_at=user.created_at,
  47. )
  48. )
  49. session.commit()
  50. def get_admin_user_by_username(self, username: str) -> AdminUser | None:
  51. with self._session_factory() as session:
  52. record = session.scalar(
  53. select(AdminUserRecord).where(AdminUserRecord.username == username)
  54. )
  55. return self._to_admin_user(record) if record else None
  56. def get_admin_user(self, user_id: str) -> AdminUser | None:
  57. with self._session_factory() as session:
  58. record = session.get(AdminUserRecord, user_id)
  59. return self._to_admin_user(record) if record else None
  60. def list_admin_users(self) -> list[AdminUser]:
  61. with self._session_factory() as session:
  62. records = session.scalars(
  63. select(AdminUserRecord).order_by(AdminUserRecord.username)
  64. ).all()
  65. return [self._to_admin_user(record) for record in records]
  66. def save_admin_user(self, user: AdminUser) -> None:
  67. with self._session_factory() as session:
  68. roles: list[RoleRecord] = []
  69. for role_code in user.roles:
  70. role = session.scalar(select(RoleRecord).where(RoleRecord.code == role_code))
  71. if role is None:
  72. role = RoleRecord(
  73. id=new_ulid(),
  74. code=role_code,
  75. name=role_code,
  76. data_scope=user.data_scope,
  77. status="ACTIVE",
  78. )
  79. session.add(role)
  80. permissions: list[PermissionRecord] = []
  81. for permission_code in user.permissions:
  82. permission = session.scalar(
  83. select(PermissionRecord).where(PermissionRecord.code == permission_code)
  84. )
  85. if permission is None:
  86. resource, _, action = permission_code.partition(":")
  87. permission = PermissionRecord(
  88. id=new_ulid(),
  89. code=permission_code,
  90. name=permission_code,
  91. resource=resource,
  92. action=action or "read",
  93. )
  94. session.add(permission)
  95. permissions.append(permission)
  96. role.permissions = permissions
  97. roles.append(role)
  98. record = session.get(AdminUserRecord, user.id)
  99. if record is None:
  100. record = AdminUserRecord(
  101. id=user.id,
  102. username=user.username,
  103. password_hash=user.password_hash,
  104. display_name=user.display_name,
  105. status=user.status,
  106. )
  107. session.add(record)
  108. else:
  109. record.username = user.username
  110. record.password_hash = user.password_hash
  111. record.display_name = user.display_name
  112. record.status = user.status
  113. record.roles = roles
  114. session.commit()
  115. def list_roles(self) -> list[AdminRole]:
  116. with self._session_factory() as session:
  117. records = session.scalars(select(RoleRecord).order_by(RoleRecord.code)).all()
  118. return [
  119. AdminRole(
  120. code=record.code,
  121. name=record.name,
  122. data_scope=record.data_scope,
  123. status=record.status,
  124. permissions=tuple(sorted(permission.code for permission in record.permissions)),
  125. )
  126. for record in records
  127. ]
  128. def save_role(self, role: AdminRole) -> None:
  129. with self._session_factory() as session:
  130. record = session.scalar(select(RoleRecord).where(RoleRecord.code == role.code))
  131. if record is None:
  132. raise KeyError(role.code)
  133. permissions: list[PermissionRecord] = []
  134. for permission_code in role.permissions:
  135. permission = session.scalar(
  136. select(PermissionRecord).where(PermissionRecord.code == permission_code)
  137. )
  138. if permission is None:
  139. resource, _, action = permission_code.partition(":")
  140. permission = PermissionRecord(
  141. id=new_ulid(),
  142. code=permission_code,
  143. name=permission_code,
  144. resource=resource,
  145. action=action or "read",
  146. )
  147. session.add(permission)
  148. permissions.append(permission)
  149. record.name = role.name
  150. record.data_scope = role.data_scope
  151. record.status = role.status
  152. record.permissions = permissions
  153. session.commit()
  154. def get_session(self, session_id: str) -> AuthSession | None:
  155. with self._session_factory() as session:
  156. record = session.get(AuthSessionRecord, session_id)
  157. if record is None:
  158. return None
  159. return AuthSession(
  160. id=record.id,
  161. subject_type=record.subject_type,
  162. subject_id=record.subject_id,
  163. refresh_jti_hash=record.refresh_jti_hash,
  164. expires_at=record.expires_at,
  165. revoked_at=record.revoked_at,
  166. )
  167. def save_session(self, auth_session: AuthSession) -> None:
  168. with self._session_factory() as session:
  169. session.merge(
  170. AuthSessionRecord(
  171. id=auth_session.id,
  172. subject_type=auth_session.subject_type,
  173. subject_id=auth_session.subject_id,
  174. refresh_jti_hash=auth_session.refresh_jti_hash,
  175. expires_at=auth_session.expires_at,
  176. revoked_at=auth_session.revoked_at,
  177. )
  178. )
  179. session.commit()
  180. @staticmethod
  181. def _to_h5_user(record: H5UserRecord) -> H5User:
  182. return H5User(
  183. id=record.id,
  184. mobile=record.mobile,
  185. mobile_masked=record.mobile_masked,
  186. display_name=record.display_name,
  187. status=record.status,
  188. created_at=record.created_at,
  189. )
  190. @staticmethod
  191. def _to_admin_user(record: AdminUserRecord) -> AdminUser:
  192. permissions = {permission.code for role in record.roles for permission in role.permissions}
  193. scopes = {role.data_scope for role in record.roles}
  194. data_scope = "ALL" if "ALL" in scopes else sorted(scopes)[0]
  195. return AdminUser(
  196. id=record.id,
  197. username=record.username,
  198. password_hash=record.password_hash,
  199. display_name=record.display_name,
  200. status=record.status,
  201. roles=tuple(sorted(role.code for role in record.roles)),
  202. permissions=tuple(sorted(permissions)),
  203. data_scope=data_scope,
  204. )
  205. class SqlAlchemyCatalogRepository(CatalogRepository):
  206. def __init__(self, session_factory: SessionFactory) -> None:
  207. self._session_factory = session_factory
  208. def save_product(self, product: Product) -> None:
  209. with self._session_factory() as session:
  210. session.merge(
  211. ProductRecord(
  212. id=product.id,
  213. product_code=product.product_code,
  214. name=product.name,
  215. category=product.category,
  216. summary=product.summary,
  217. status=product.status,
  218. )
  219. )
  220. session.commit()
  221. def save_version(self, version: ProductVersion) -> None:
  222. with self._session_factory() as session:
  223. record = session.get(ProductVersionRecord, version.id)
  224. if record is None:
  225. record = ProductVersionRecord(
  226. id=version.id,
  227. product_id=version.product_id,
  228. version_no=version.version_no,
  229. )
  230. session.add(record)
  231. record.version_no = version.version_no
  232. record.status = version.status
  233. record.effective_from = version.effective_from
  234. record.effective_to = version.effective_to
  235. record.rule_version = version.rule_version
  236. record.rate_version = version.rate_version
  237. record.terms_summary = version.terms_summary
  238. record.plans = [
  239. PlanRecord(
  240. id=plan.id,
  241. plan_code=plan.code,
  242. name=plan.name,
  243. summary=plan.summary,
  244. status=plan.status,
  245. premium_cents=plan.premium_cents,
  246. coverage_amount_cents=plan.coverage_amount_cents,
  247. min_age=plan.min_age,
  248. max_age=plan.max_age,
  249. )
  250. for plan in version.plans
  251. ]
  252. session.commit()
  253. def delete_version(self, version_id: str) -> None:
  254. with self._session_factory() as session:
  255. record = session.get(ProductVersionRecord, version_id)
  256. if record is not None:
  257. session.delete(record)
  258. session.commit()
  259. def save_change_log(self, log: ProductChangeLog) -> None:
  260. with self._session_factory() as session:
  261. session.merge(
  262. ProductChangeLogRecord(
  263. id=log.id,
  264. product_id=log.product_id,
  265. version_id=log.version_id,
  266. action=log.action,
  267. actor_id=log.actor_id,
  268. actor_name=log.actor_name,
  269. detail_json=log.detail,
  270. created_at=log.created_at,
  271. )
  272. )
  273. session.commit()
  274. def list_change_logs(self, product_id: str) -> list[ProductChangeLog]:
  275. with self._session_factory() as session:
  276. records = session.scalars(
  277. select(ProductChangeLogRecord)
  278. .where(ProductChangeLogRecord.product_id == product_id)
  279. .order_by(ProductChangeLogRecord.created_at.desc())
  280. ).all()
  281. return [
  282. ProductChangeLog(
  283. id=record.id,
  284. product_id=record.product_id,
  285. version_id=record.version_id,
  286. action=record.action,
  287. actor_id=record.actor_id,
  288. actor_name=record.actor_name,
  289. detail=dict(record.detail_json),
  290. created_at=record.created_at,
  291. )
  292. for record in records
  293. ]
  294. def list_products(self) -> list[Product]:
  295. with self._session_factory() as session:
  296. records = session.scalars(select(ProductRecord)).all()
  297. return [
  298. Product(
  299. id=record.id,
  300. product_code=record.product_code,
  301. name=record.name,
  302. category=record.category,
  303. summary=record.summary,
  304. status=record.status,
  305. )
  306. for record in records
  307. ]
  308. def list_versions(self, product_id: str) -> list[ProductVersion]:
  309. with self._session_factory() as session:
  310. records = session.scalars(
  311. select(ProductVersionRecord).where(ProductVersionRecord.product_id == product_id)
  312. ).all()
  313. return [
  314. ProductVersion(
  315. id=record.id,
  316. product_id=record.product_id,
  317. version_no=record.version_no,
  318. status=record.status,
  319. effective_from=record.effective_from,
  320. effective_to=record.effective_to,
  321. rule_version=record.rule_version,
  322. rate_version=record.rate_version,
  323. terms_summary=record.terms_summary,
  324. plans=tuple(
  325. Plan(
  326. id=plan.id,
  327. code=plan.plan_code,
  328. name=plan.name,
  329. summary=plan.summary,
  330. status=plan.status,
  331. premium_cents=plan.premium_cents,
  332. coverage_amount_cents=plan.coverage_amount_cents,
  333. min_age=plan.min_age,
  334. max_age=plan.max_age,
  335. )
  336. for plan in record.plans
  337. ),
  338. )
  339. for record in records
  340. ]