s.Middleware.go 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200
  1. package jwtx
  2. import (
  3. "context"
  4. "crypto/ed25519"
  5. "errors"
  6. "time"
  7. "git.ooo.ink/root/go-kit/src/gin/jwtx/db/dao/model"
  8. "git.ooo.ink/root/go-kit/src/gin/jwtx/tool"
  9. "git.ooo.ink/root/go-kit/src/logx"
  10. "github.com/gin-gonic/gin"
  11. "github.com/golang-jwt/jwt/v5"
  12. )
  13. // 中间件认证失败后的处理函数 ..
  14. var MiddlewareFailHandlerFunc = func(c *gin.Context, err error) {
  15. logx.Debug().Err(err)
  16. c.AbortWithStatus(401)
  17. }
  18. // jwtx token 中间件(token 解析、校验、刷新) ..
  19. //
  20. // requestGroup string 请求访问的分组
  21. //
  22. // e.g.
  23. //
  24. // jwtx.Singleton.Middleware("admin")
  25. // jwtx.Singleton.Middleware("merchant")
  26. func (s *SingletonT) Middleware(requestGroup string) gin.HandlerFunc {
  27. return func(c *gin.Context) {
  28. // 解析 token
  29. // 从 Authorization 头中提取 JWT token 并进行验证
  30. tokenStr := c.GetHeader("Authorization")
  31. if tokenStr == "" {
  32. MiddlewareFailHandlerFunc(c, errors.New("authorization header is required"))
  33. return
  34. }
  35. // 使用对应分组的公钥验证 token 签名
  36. claims, err := tool.ParseToken(tokenStr, s.PrivateKey[requestGroup].Public().(*ed25519.PublicKey))
  37. if err != nil {
  38. MiddlewareFailHandlerFunc(c, err)
  39. return
  40. }
  41. // 提取 claims
  42. var (
  43. now = time.Now() // 当前时间
  44. iat, _ = claims.GetIssuedAt() // 签发时间
  45. exp, _ = claims.GetExpirationTime() // 过期时间
  46. tid = uint64(claims["tid"].(float64)) // token ID
  47. )
  48. // token 过期时间校验
  49. if exp.Unix() < now.Unix() {
  50. MiddlewareFailHandlerFunc(c, errors.New("token has expired"))
  51. return
  52. }
  53. // 数据库中查找 token
  54. token, err := s.getDBToken(requestGroup, tid)
  55. if err != nil {
  56. MiddlewareFailHandlerFunc(c, err)
  57. return
  58. }
  59. // 校验数据库 token 信息
  60. err = s.checkDBToken(requestGroup, token, now, c.ClientIP())
  61. if err != nil {
  62. MiddlewareFailHandlerFunc(c, err)
  63. return
  64. }
  65. // 刷新 token
  66. newToken, err := s.refreshToken(requestGroup, token, now, iat)
  67. if err != nil {
  68. MiddlewareFailHandlerFunc(c, err)
  69. return
  70. }
  71. // 将当前登录的账户信息存储上下文
  72. c.Set("jwtx.accountID", token.AccountID) // 账户 ID
  73. c.Set("jwtx.loginGroup", token.LoginGroup) // 登录的分组
  74. c.Set("jwtx.loginTerminal", token.LoginTerminal) // 登录的终端
  75. c.Set("jwtx.makeTokenIP", token.MakeTokenIP) // 首次请求生成 token 的 IP 地址
  76. // 请求前
  77. c.Next()
  78. // 请求后
  79. // 响应头填充刷新后的 token
  80. c.Writer.Header().Add("token", newToken)
  81. }
  82. }
  83. // 获取数据库 token 信息
  84. //
  85. // 参数说明:
  86. // - group string: token 分组标识,用于获取对应的数据库连接
  87. // - tid uint64: token 记录的唯一标识符
  88. //
  89. // 返回值:
  90. // - *model.JwtxToken: 查询到的 token 记录
  91. // - error: 查询过程中出现的错误
  92. //
  93. // 功能说明:
  94. //
  95. // 根据 token ID 从数据库中查询对应的 token 记录
  96. // 使用标准的 GORM 查询方式,避免使用 query 包
  97. func (s *SingletonT) getDBToken(group string, tid uint64) (token *model.JwtxToken, err error) {
  98. // 使用标准的 GORM 查询方式查找 token
  99. // 根据 token ID 查询对应的记录
  100. token = &model.JwtxToken{}
  101. err = s.DB[group].WithContext(context.Background()).
  102. Where("id = ?", tid).
  103. First(token).
  104. Error
  105. if err != nil {
  106. // 查询失败,返回错误
  107. return nil, err
  108. }
  109. // 查询成功,返回 token 记录
  110. return token, nil
  111. }
  112. // 校验数据库 token 信息
  113. func (s *SingletonT) checkDBToken(group string, token *model.JwtxToken, now time.Time, clientIP string) (err error) {
  114. // 分组校验
  115. if token.LoginGroup != group {
  116. return errors.New("auth group fail")
  117. }
  118. // 过期时间校验
  119. if token.ExpirationAt.Unix() < now.Unix() {
  120. return errors.New("the token has expired")
  121. }
  122. // IP 一致性校验
  123. if s.Config[group].CheckIP {
  124. if clientIP != token.MakeTokenIP {
  125. return errors.New("client ip is changed, please login again")
  126. }
  127. }
  128. return
  129. }
  130. // 自动刷新 token
  131. func (s *SingletonT) refreshToken(group string, token *model.JwtxToken, now time.Time, iat *jwt.NumericDate) (newToken string, err error) {
  132. var config = s.Config[group]
  133. if iat.Unix() == token.FinalRefreshAt.Unix() { // token 未刷新
  134. // 原始的 token 过期时间
  135. expTime := token.ExpirationAt
  136. // 需要刷新 token
  137. if iat.Unix()+config.RefreshInterval < now.Unix() {
  138. // 自动续期
  139. if config.AutomaticRenewal {
  140. expTime = now.Add(time.Duration(config.AccessExpireByHour) * time.Hour)
  141. }
  142. // 构造 token 字符串(过期时间不变,签发时间顺延)
  143. newToken, err = tool.GenerateToken(s.PrivateKey[group], jwt.MapClaims{
  144. "iat": now.Unix(), // 签发时间
  145. "exp": expTime.Unix(), // 过期时间
  146. "tid": token.ID, // jwt token ID
  147. })
  148. if err != nil {
  149. return "", err
  150. }
  151. // 更新数据库 - 使用标准的 GORM 更新方式
  152. // 更新 token 的刷新时间信息
  153. err = s.DB[group].WithContext(context.Background()).
  154. Model(&model.JwtxToken{}).
  155. Where("id = ?", token.ID).
  156. Updates(map[string]interface{}{
  157. "last_refresh_at": token.FinalRefreshAt, // 上次刷新时间设置为之前的最后刷新时间
  158. "final_refresh_at": now, // 最后刷新时间更新为当前时间
  159. }).Error
  160. if err != nil {
  161. return "", err
  162. }
  163. }
  164. } else if iat.Unix() == token.LastRefreshAt.Unix() { // token 已刷新
  165. // 当前时间 超出 并发容错时间(不允许继续使用)
  166. if now.Unix() > token.FinalRefreshAt.Unix()+config.FaultTolerance {
  167. return "", errors.New("out of concurrent fault tolerance time")
  168. }
  169. }
  170. return
  171. }