s.Middleware.go 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170
  1. package jwtx
  2. import (
  3. "context"
  4. "errors"
  5. "time"
  6. "github.com/5-say/go-tool/gin/jwtx/db/dao/model"
  7. "github.com/5-say/go-tool/gin/jwtx/db/dao/query"
  8. "github.com/5-say/go-tool/gin/jwtx/tool"
  9. "github.com/5-say/go-tool/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. claims, err := tool.ParseToken(c.GetHeader("Authorization"), &s.PrivateKey[requestGroup].PublicKey)
  30. if err != nil {
  31. MiddlewareFailHandlerFunc(c, err)
  32. return
  33. }
  34. // 提取 claims
  35. var (
  36. now = time.Now() // 当前时间
  37. iat, _ = claims.GetIssuedAt() // 签发时间
  38. exp, _ = claims.GetExpirationTime() // 过期时间
  39. tid = uint64(claims["tid"].(float64)) // token ID
  40. )
  41. // token 过期时间校验
  42. if exp.Unix() < now.Unix() {
  43. MiddlewareFailHandlerFunc(c, errors.New("token has expired"))
  44. return
  45. }
  46. // 数据库中查找 token
  47. token, err := s.getDBToken(requestGroup, tid)
  48. if err != nil {
  49. MiddlewareFailHandlerFunc(c, err)
  50. return
  51. }
  52. // 校验数据库 token 信息
  53. err = s.checkDBToken(requestGroup, token, now, c.ClientIP())
  54. if err != nil {
  55. MiddlewareFailHandlerFunc(c, err)
  56. return
  57. }
  58. // 刷新 token
  59. newToken, err := s.refreshToken(requestGroup, token, now, iat)
  60. if err != nil {
  61. MiddlewareFailHandlerFunc(c, err)
  62. return
  63. }
  64. // 将当前登录的账户信息存储上下文
  65. c.Set("jwtx.accountID", token.AccountID) // 账户 ID
  66. c.Set("jwtx.loginGroup", token.LoginGroup) // 登录的分组
  67. c.Set("jwtx.loginTerminal", token.LoginTerminal) // 登录的终端
  68. c.Set("jwtx.makeTokenIP", token.MakeTokenIP) // 首次请求生成 token 的 IP 地址
  69. // 请求前
  70. c.Next()
  71. // 请求后
  72. // 响应头填充刷新后的 token
  73. c.Writer.Header().Add("token", newToken)
  74. }
  75. }
  76. // 获取数据库 token 信息
  77. func (s *SingletonT) getDBToken(group string, tid uint64) (token *model.JwtxToken, err error) {
  78. // 查找 token
  79. q := query.Use(s.DB[group])
  80. o := q.JwtxToken
  81. token, err = o.WithContext(context.Background()).Where(o.ID.Eq(tid)).First()
  82. if err != nil {
  83. return nil, err
  84. }
  85. return token, nil
  86. }
  87. // 校验数据库 token 信息
  88. func (s *SingletonT) checkDBToken(group string, token *model.JwtxToken, now time.Time, clientIP string) (err error) {
  89. // 分组校验
  90. if token.LoginGroup != group {
  91. return errors.New("auth group fail")
  92. }
  93. // 过期时间校验
  94. if token.ExpirationAt.Unix() < now.Unix() {
  95. return errors.New("the token has expired")
  96. }
  97. // IP 一致性校验
  98. if s.Config[group].CheckIP {
  99. if clientIP != token.MakeTokenIP {
  100. return errors.New("client ip is changed, please login again")
  101. }
  102. }
  103. return
  104. }
  105. // 自动刷新 token
  106. func (s *SingletonT) refreshToken(group string, token *model.JwtxToken, now time.Time, iat *jwt.NumericDate) (newToken string, err error) {
  107. var config = s.Config[group]
  108. if iat.Unix() == token.FinalRefreshAt.Unix() { // token 未刷新
  109. // 原始的 token 过期时间
  110. expTime := token.ExpirationAt
  111. // 需要刷新 token
  112. if iat.Unix()+config.RefreshInterval < now.Unix() {
  113. // 自动续期
  114. if config.AutomaticRenewal {
  115. expTime = now.Add(time.Duration(config.AccessExpireByHour) * time.Hour)
  116. }
  117. // 构造 token 字符串(过期时间不变,签发时间顺延)
  118. newToken, err = tool.GenerateToken(s.PrivateKey[group], jwt.MapClaims{
  119. "iat": now.Unix(), // 签发时间
  120. "exp": expTime.Unix(), // 过期时间
  121. "tid": token.ID, // jwt token ID
  122. }, tool.P_384)
  123. if err != nil {
  124. return "", err
  125. }
  126. // 更新数据库
  127. q := query.Use(s.DB[group])
  128. o := q.JwtxToken
  129. _, err = o.WithContext(context.Background()).Where(o.ID.Eq(token.ID)).UpdateSimple(
  130. o.LastRefreshAt.Value(token.FinalRefreshAt),
  131. o.FinalRefreshAt.Value(now),
  132. )
  133. if err != nil {
  134. return "", err
  135. }
  136. }
  137. } else if iat.Unix() == token.LastRefreshAt.Unix() { // token 已刷新
  138. // 当前时间 超出 并发容错时间(不允许继续使用)
  139. if now.Unix() > token.FinalRefreshAt.Unix()+config.FaultTolerance {
  140. return "", errors.New("out of concurrent fault tolerance time")
  141. }
  142. }
  143. return
  144. }