repository.py 4.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115
  1. from dataclasses import replace
  2. from typing import Protocol
  3. from app.domains.identity.models import AdminRole, AdminUser, AuthSession, H5User
  4. class IdentityRepository(Protocol):
  5. def get_h5_user_by_mobile(self, mobile: str) -> H5User | None: ...
  6. def get_h5_user(self, user_id: str) -> H5User | None: ...
  7. def list_h5_users(self) -> list[H5User]: ...
  8. def save_h5_user(self, user: H5User) -> None: ...
  9. def get_admin_user_by_username(self, username: str) -> AdminUser | None: ...
  10. def get_admin_user(self, user_id: str) -> AdminUser | None: ...
  11. def list_admin_users(self) -> list[AdminUser]: ...
  12. def save_admin_user(self, user: AdminUser) -> None: ...
  13. def list_roles(self) -> list[AdminRole]: ...
  14. def save_role(self, role: AdminRole) -> None: ...
  15. def get_session(self, session_id: str) -> AuthSession | None: ...
  16. def save_session(self, session: AuthSession) -> None: ...
  17. class InMemoryIdentityRepository:
  18. def __init__(self) -> None:
  19. self._h5_users: dict[str, H5User] = {}
  20. self._h5_user_ids_by_mobile: dict[str, str] = {}
  21. self._admin_users: dict[str, AdminUser] = {}
  22. self._admin_user_ids_by_username: dict[str, str] = {}
  23. self._roles: dict[str, AdminRole] = {}
  24. self._sessions: dict[str, AuthSession] = {}
  25. def get_h5_user_by_mobile(self, mobile: str) -> H5User | None:
  26. user_id = self._h5_user_ids_by_mobile.get(mobile)
  27. return self._h5_users.get(user_id) if user_id else None
  28. def get_h5_user(self, user_id: str) -> H5User | None:
  29. return self._h5_users.get(user_id)
  30. def list_h5_users(self) -> list[H5User]:
  31. return sorted(
  32. self._h5_users.values(),
  33. key=lambda user: user.created_at,
  34. reverse=True,
  35. )
  36. def save_h5_user(self, user: H5User) -> None:
  37. self._h5_users[user.id] = user
  38. self._h5_user_ids_by_mobile[user.mobile] = user.id
  39. def get_admin_user_by_username(self, username: str) -> AdminUser | None:
  40. user_id = self._admin_user_ids_by_username.get(username)
  41. return self._admin_users.get(user_id) if user_id else None
  42. def get_admin_user(self, user_id: str) -> AdminUser | None:
  43. return self._admin_users.get(user_id)
  44. def list_admin_users(self) -> list[AdminUser]:
  45. return sorted(self._admin_users.values(), key=lambda user: user.username)
  46. def save_admin_user(self, user: AdminUser) -> None:
  47. self._admin_users[user.id] = user
  48. self._admin_user_ids_by_username[user.username] = user.id
  49. for role_code in user.roles:
  50. self._roles.setdefault(
  51. role_code,
  52. AdminRole(
  53. code=role_code,
  54. name=role_code,
  55. data_scope=user.data_scope,
  56. status="ACTIVE",
  57. permissions=user.permissions,
  58. ),
  59. )
  60. def list_roles(self) -> list[AdminRole]:
  61. return sorted(self._roles.values(), key=lambda role: role.code)
  62. def save_role(self, role: AdminRole) -> None:
  63. if role.code not in self._roles:
  64. raise KeyError(role.code)
  65. self._roles[role.code] = role
  66. for user_id, user in tuple(self._admin_users.items()):
  67. if role.code in user.roles:
  68. assigned_roles = [self._roles[code] for code in user.roles]
  69. permissions = tuple(
  70. sorted(
  71. {
  72. permission
  73. for assigned_role in assigned_roles
  74. for permission in assigned_role.permissions
  75. }
  76. )
  77. )
  78. scopes = {assigned_role.data_scope for assigned_role in assigned_roles}
  79. self._admin_users[user_id] = replace(
  80. user,
  81. permissions=permissions,
  82. data_scope="ALL" if "ALL" in scopes else sorted(scopes)[0],
  83. )
  84. def get_session(self, session_id: str) -> AuthSession | None:
  85. return self._sessions.get(session_id)
  86. def save_session(self, session: AuthSession) -> None:
  87. self._sessions[session.id] = session