sysUserModel.go 9.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230
  1. package user
  2. import (
  3. "context"
  4. "database/sql"
  5. "errors"
  6. "fmt"
  7. "strings"
  8. "time"
  9. "github.com/zeromicro/go-zero/core/stores/cache"
  10. "github.com/zeromicro/go-zero/core/stores/sqlx"
  11. )
  12. var ErrUpdateConflict = errors.New("update conflict: data has been modified by another operation")
  13. // ErrTokenVersionMismatch 表示令牌版本与数据库当前版本不一致,刷新令牌失败。
  14. // 典型场景:refreshToken rotation 并发到达 —— 只有持有当前 tokenVersion 的那一次能原子递增成功,
  15. // 其余全部返回该错误,防止两个请求都"换到"新令牌(导致会话劫持)。
  16. var ErrTokenVersionMismatch = errors.New("token version mismatch")
  17. var _ SysUserModel = (*customSysUserModel)(nil)
  18. type (
  19. SysUserModel interface {
  20. sysUserModel
  21. FindListByPage(ctx context.Context, page, pageSize int64) ([]*SysUser, int64, error)
  22. FindListByProductMembers(ctx context.Context, productCode string, page, pageSize int64) ([]*SysUser, map[int64]string, int64, error)
  23. FindByIds(ctx context.Context, ids []int64) ([]*SysUser, error)
  24. FindIdsByDeptId(ctx context.Context, deptId int64) ([]int64, error)
  25. UpdateProfile(ctx context.Context, id int64, username string, nickname, email, phone, remark string, deptId, newStatus int64, statusChanged bool, expectedUpdateTime int64) error
  26. UpdatePassword(ctx context.Context, id int64, password string, mustChangePassword int64) error
  27. UpdateStatus(ctx context.Context, id int64, status int64) error
  28. IncrementTokenVersion(ctx context.Context, id int64) (int64, error)
  29. IncrementTokenVersionIfMatch(ctx context.Context, id int64, username string, expected int64) (int64, error)
  30. }
  31. customSysUserModel struct {
  32. *defaultSysUserModel
  33. }
  34. )
  35. func NewSysUserModel(conn sqlx.SqlConn, c cache.CacheConf, cachePrefix string, opts ...cache.Option) SysUserModel {
  36. return &customSysUserModel{
  37. defaultSysUserModel: newSysUserModel(conn, c, cachePrefix, opts...),
  38. }
  39. }
  40. func (m *customSysUserModel) FindListByPage(ctx context.Context, page, pageSize int64) ([]*SysUser, int64, error) {
  41. var total int64
  42. countQuery := fmt.Sprintf("SELECT COUNT(*) FROM %s", m.table)
  43. if err := m.QueryRowNoCacheCtx(ctx, &total, countQuery); err != nil {
  44. return nil, 0, err
  45. }
  46. var list []*SysUser
  47. query := fmt.Sprintf("SELECT %s FROM %s ORDER BY id DESC LIMIT ?,?", sysUserRows, m.table)
  48. if err := m.QueryRowsNoCacheCtx(ctx, &list, query, (page-1)*pageSize, pageSize); err != nil {
  49. return nil, 0, err
  50. }
  51. return list, total, nil
  52. }
  53. type UserWithMemberType struct {
  54. SysUser
  55. MemberType string `db:"memberType"`
  56. }
  57. func (m *customSysUserModel) FindListByProductMembers(ctx context.Context, productCode string, page, pageSize int64) ([]*SysUser, map[int64]string, int64, error) {
  58. memberTable := "`sys_product_member`"
  59. var total int64
  60. countQuery := fmt.Sprintf("SELECT COUNT(*) FROM %s u INNER JOIN %s pm ON u.`id` = pm.`userId` WHERE pm.`productCode` = ?", m.table, memberTable)
  61. if err := m.QueryRowNoCacheCtx(ctx, &total, countQuery, productCode); err != nil {
  62. return nil, nil, 0, err
  63. }
  64. var list []*UserWithMemberType
  65. fields := strings.Join(sysUserFieldNames, ",u.")
  66. query := fmt.Sprintf("SELECT u.%s, pm.`memberType` FROM %s u INNER JOIN %s pm ON u.`id` = pm.`userId` WHERE pm.`productCode` = ? ORDER BY u.`id` DESC LIMIT ?,?", fields, m.table, memberTable)
  67. if err := m.QueryRowsNoCacheCtx(ctx, &list, query, productCode, (page-1)*pageSize, pageSize); err != nil {
  68. return nil, nil, 0, err
  69. }
  70. users := make([]*SysUser, len(list))
  71. memberMap := make(map[int64]string, len(list))
  72. for i, item := range list {
  73. users[i] = &item.SysUser
  74. memberMap[item.Id] = item.MemberType
  75. }
  76. return users, memberMap, total, nil
  77. }
  78. func (m *customSysUserModel) FindIdsByDeptId(ctx context.Context, deptId int64) ([]int64, error) {
  79. var ids []int64
  80. query := fmt.Sprintf("SELECT `id` FROM %s WHERE `deptId` = ?", m.table)
  81. if err := m.QueryRowsNoCacheCtx(ctx, &ids, query, deptId); err != nil {
  82. return nil, err
  83. }
  84. return ids, nil
  85. }
  86. func (m *customSysUserModel) UpdateProfile(ctx context.Context, id int64, username string, nickname, email, phone, remark string, deptId, newStatus int64, statusChanged bool, expectedUpdateTime int64) error {
  87. sysUserIdKey := fmt.Sprintf("%s%v", cacheSysUserIdPrefix, id)
  88. sysUserUsernameKey := fmt.Sprintf("%s%v", cacheSysUserUsernamePrefix, username)
  89. now := time.Now().Unix()
  90. res, err := m.ExecCtx(ctx, func(ctx context.Context, conn sqlx.SqlConn) (sql.Result, error) {
  91. if statusChanged {
  92. query := fmt.Sprintf("UPDATE %s SET `nickname`=?, `email`=?, `phone`=?, `remark`=?, `deptId`=?, `status`=?, `tokenVersion`=`tokenVersion`+1, `updateTime`=? WHERE `id`=? AND `updateTime`=?", m.table)
  93. return conn.ExecCtx(ctx, query, nickname, email, phone, remark, deptId, newStatus, now, id, expectedUpdateTime)
  94. }
  95. query := fmt.Sprintf("UPDATE %s SET `nickname`=?, `email`=?, `phone`=?, `remark`=?, `deptId`=?, `updateTime`=? WHERE `id`=? AND `updateTime`=?", m.table)
  96. return conn.ExecCtx(ctx, query, nickname, email, phone, remark, deptId, now, id, expectedUpdateTime)
  97. }, sysUserIdKey, sysUserUsernameKey)
  98. if err != nil {
  99. return err
  100. }
  101. affected, _ := res.RowsAffected()
  102. if affected == 0 {
  103. return ErrUpdateConflict
  104. }
  105. return nil
  106. }
  107. func (m *customSysUserModel) UpdatePassword(ctx context.Context, id int64, password string, mustChangePassword int64) error {
  108. data, err := m.FindOne(ctx, id)
  109. if err != nil {
  110. return err
  111. }
  112. sysUserIdKey := fmt.Sprintf("%s%v", cacheSysUserIdPrefix, id)
  113. sysUserUsernameKey := fmt.Sprintf("%s%v", cacheSysUserUsernamePrefix, data.Username)
  114. _, err = m.ExecCtx(ctx, func(ctx context.Context, conn sqlx.SqlConn) (sql.Result, error) {
  115. query := fmt.Sprintf("UPDATE %s SET `password` = ?, `mustChangePassword` = ?, `tokenVersion` = `tokenVersion` + 1, `updateTime` = ? WHERE `id` = ?", m.table)
  116. return conn.ExecCtx(ctx, query, password, mustChangePassword, time.Now().Unix(), id)
  117. }, sysUserIdKey, sysUserUsernameKey)
  118. return err
  119. }
  120. func (m *customSysUserModel) UpdateStatus(ctx context.Context, id int64, status int64) error {
  121. data, err := m.FindOne(ctx, id)
  122. if err != nil {
  123. return err
  124. }
  125. sysUserIdKey := fmt.Sprintf("%s%v", cacheSysUserIdPrefix, id)
  126. sysUserUsernameKey := fmt.Sprintf("%s%v", cacheSysUserUsernamePrefix, data.Username)
  127. _, err = m.ExecCtx(ctx, func(ctx context.Context, conn sqlx.SqlConn) (sql.Result, error) {
  128. query := fmt.Sprintf("UPDATE %s SET `status` = ?, `tokenVersion` = `tokenVersion` + 1, `updateTime` = ? WHERE `id` = ?", m.table)
  129. return conn.ExecCtx(ctx, query, status, time.Now().Unix(), id)
  130. }, sysUserIdKey, sysUserUsernameKey)
  131. return err
  132. }
  133. func (m *customSysUserModel) IncrementTokenVersion(ctx context.Context, id int64) (int64, error) {
  134. data, err := m.FindOne(ctx, id)
  135. if err != nil {
  136. return 0, err
  137. }
  138. sysUserIdKey := fmt.Sprintf("%s%v", cacheSysUserIdPrefix, id)
  139. sysUserUsernameKey := fmt.Sprintf("%s%v", cacheSysUserUsernamePrefix, data.Username)
  140. var newVersion int64
  141. err = m.TransactCtx(ctx, func(ctx context.Context, session sqlx.Session) error {
  142. query := fmt.Sprintf("UPDATE %s SET `tokenVersion` = LAST_INSERT_ID(`tokenVersion` + 1), `updateTime` = ? WHERE `id` = ?", m.table)
  143. if _, err := session.ExecCtx(ctx, query, time.Now().Unix(), id); err != nil {
  144. return err
  145. }
  146. return session.QueryRowCtx(ctx, &newVersion, "SELECT LAST_INSERT_ID()")
  147. })
  148. if err != nil {
  149. return 0, err
  150. }
  151. _ = m.DelCacheCtx(ctx, sysUserIdKey, sysUserUsernameKey)
  152. return newVersion, nil
  153. }
  154. // IncrementTokenVersionIfMatch 原子递增 tokenVersion;仅当 DB 里当前 tokenVersion == expected 时才会生效。
  155. // 这是 refreshToken rotation 的原子 CAS:两个并发的刷新请求只有一个能命中 WHERE tokenVersion=expected,
  156. // 另一个 affected=0 返回 ErrTokenVersionMismatch,从而避免"两边都换到新令牌"的会话劫持窗口。
  157. //
  158. // 由上游透传 username 以便构造 cacheSysUserUsernamePrefix 的缓存键进行失效,避免为此多查一次 FindOne
  159. // (见审计 M-8)。上游通常已经通过 UserDetailsLoader.Load 拿到 username,零额外成本。
  160. func (m *customSysUserModel) IncrementTokenVersionIfMatch(ctx context.Context, id int64, username string, expected int64) (int64, error) {
  161. sysUserIdKey := fmt.Sprintf("%s%v", cacheSysUserIdPrefix, id)
  162. sysUserUsernameKey := fmt.Sprintf("%s%v", cacheSysUserUsernamePrefix, username)
  163. var newVersion int64
  164. err := m.TransactCtx(ctx, func(ctx context.Context, session sqlx.Session) error {
  165. query := fmt.Sprintf("UPDATE %s SET `tokenVersion` = LAST_INSERT_ID(`tokenVersion` + 1), `updateTime` = ? WHERE `id` = ? AND `tokenVersion` = ?", m.table)
  166. res, err := session.ExecCtx(ctx, query, time.Now().Unix(), id, expected)
  167. if err != nil {
  168. return err
  169. }
  170. affected, _ := res.RowsAffected()
  171. if affected == 0 {
  172. return ErrTokenVersionMismatch
  173. }
  174. return session.QueryRowCtx(ctx, &newVersion, "SELECT LAST_INSERT_ID()")
  175. })
  176. if err != nil {
  177. return 0, err
  178. }
  179. _ = m.DelCacheCtx(ctx, sysUserIdKey, sysUserUsernameKey)
  180. return newVersion, nil
  181. }
  182. func (m *customSysUserModel) FindByIds(ctx context.Context, ids []int64) ([]*SysUser, error) {
  183. if len(ids) == 0 {
  184. return nil, nil
  185. }
  186. placeholders := make([]string, len(ids))
  187. args := make([]interface{}, len(ids))
  188. for i, id := range ids {
  189. placeholders[i] = "?"
  190. args[i] = id
  191. }
  192. var list []*SysUser
  193. query := fmt.Sprintf("SELECT %s FROM %s WHERE `id` IN (%s)", sysUserRows, m.table, strings.Join(placeholders, ","))
  194. if err := m.QueryRowsNoCacheCtx(ctx, &list, query, args...); err != nil {
  195. return nil, err
  196. }
  197. return list, nil
  198. }