package jwtx import ( "context" "crypto/ed25519" "errors" "time" "git.ooo.ink/root/go-kit/src/gin/jwtx/db/dao/model" "git.ooo.ink/root/go-kit/src/gin/jwtx/tool" "git.ooo.ink/root/go-kit/src/logx" "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" ) // 中间件认证失败后的处理函数 .. var MiddlewareFailHandlerFunc = func(c *gin.Context, err error) { logx.Debug().Err(err) c.AbortWithStatus(401) } // jwtx token 中间件(token 解析、校验、刷新) .. // // requestGroup string 请求访问的分组 // // e.g. // // jwtx.Singleton.Middleware("admin") // jwtx.Singleton.Middleware("merchant") func (s *SingletonT) Middleware(requestGroup string) gin.HandlerFunc { return func(c *gin.Context) { // 解析 token // 从 Authorization 头中提取 JWT token 并进行验证 tokenStr := c.GetHeader("Authorization") if tokenStr == "" { MiddlewareFailHandlerFunc(c, errors.New("authorization header is required")) return } // 使用对应分组的公钥验证 token 签名 claims, err := tool.ParseToken(tokenStr, s.PrivateKey[requestGroup].Public().(*ed25519.PublicKey)) if err != nil { MiddlewareFailHandlerFunc(c, err) return } // 提取 claims var ( now = time.Now() // 当前时间 iat, _ = claims.GetIssuedAt() // 签发时间 exp, _ = claims.GetExpirationTime() // 过期时间 tid = uint64(claims["tid"].(float64)) // token ID ) // token 过期时间校验 if exp.Unix() < now.Unix() { MiddlewareFailHandlerFunc(c, errors.New("token has expired")) return } // 数据库中查找 token token, err := s.getDBToken(requestGroup, tid) if err != nil { MiddlewareFailHandlerFunc(c, err) return } // 校验数据库 token 信息 err = s.checkDBToken(requestGroup, token, now, c.ClientIP()) if err != nil { MiddlewareFailHandlerFunc(c, err) return } // 刷新 token newToken, err := s.refreshToken(requestGroup, token, now, iat) if err != nil { MiddlewareFailHandlerFunc(c, err) return } // 将当前登录的账户信息存储上下文 c.Set("jwtx.accountID", token.AccountID) // 账户 ID c.Set("jwtx.loginGroup", token.LoginGroup) // 登录的分组 c.Set("jwtx.loginTerminal", token.LoginTerminal) // 登录的终端 c.Set("jwtx.makeTokenIP", token.MakeTokenIP) // 首次请求生成 token 的 IP 地址 // 请求前 c.Next() // 请求后 // 响应头填充刷新后的 token c.Writer.Header().Add("token", newToken) } } // 获取数据库 token 信息 // // 参数说明: // - group string: token 分组标识,用于获取对应的数据库连接 // - tid uint64: token 记录的唯一标识符 // // 返回值: // - *model.JwtxToken: 查询到的 token 记录 // - error: 查询过程中出现的错误 // // 功能说明: // // 根据 token ID 从数据库中查询对应的 token 记录 // 使用标准的 GORM 查询方式,避免使用 query 包 func (s *SingletonT) getDBToken(group string, tid uint64) (token *model.JwtxToken, err error) { // 使用标准的 GORM 查询方式查找 token // 根据 token ID 查询对应的记录 token = &model.JwtxToken{} err = s.DB[group].WithContext(context.Background()). Where("id = ?", tid). First(token). Error if err != nil { // 查询失败,返回错误 return nil, err } // 查询成功,返回 token 记录 return token, nil } // 校验数据库 token 信息 func (s *SingletonT) checkDBToken(group string, token *model.JwtxToken, now time.Time, clientIP string) (err error) { // 分组校验 if token.LoginGroup != group { return errors.New("auth group fail") } // 过期时间校验 if token.ExpirationAt.Unix() < now.Unix() { return errors.New("the token has expired") } // IP 一致性校验 if s.Config[group].CheckIP { if clientIP != token.MakeTokenIP { return errors.New("client ip is changed, please login again") } } return } // 自动刷新 token func (s *SingletonT) refreshToken(group string, token *model.JwtxToken, now time.Time, iat *jwt.NumericDate) (newToken string, err error) { var config = s.Config[group] if iat.Unix() == token.FinalRefreshAt.Unix() { // token 未刷新 // 原始的 token 过期时间 expTime := token.ExpirationAt // 需要刷新 token if iat.Unix()+config.RefreshInterval < now.Unix() { // 自动续期 if config.AutomaticRenewal { expTime = now.Add(time.Duration(config.AccessExpireByHour) * time.Hour) } // 构造 token 字符串(过期时间不变,签发时间顺延) newToken, err = tool.GenerateToken(s.PrivateKey[group], jwt.MapClaims{ "iat": now.Unix(), // 签发时间 "exp": expTime.Unix(), // 过期时间 "tid": token.ID, // jwt token ID }) if err != nil { return "", err } // 更新数据库 - 使用标准的 GORM 更新方式 // 更新 token 的刷新时间信息 err = s.DB[group].WithContext(context.Background()). Model(&model.JwtxToken{}). Where("id = ?", token.ID). Updates(map[string]interface{}{ "last_refresh_at": token.FinalRefreshAt, // 上次刷新时间设置为之前的最后刷新时间 "final_refresh_at": now, // 最后刷新时间更新为当前时间 }).Error if err != nil { return "", err } } } else if iat.Unix() == token.LastRefreshAt.Unix() { // token 已刷新 // 当前时间 超出 并发容错时间(不允许继续使用) if now.Unix() > token.FinalRefreshAt.Unix()+config.FaultTolerance { return "", errors.New("out of concurrent fault tolerance time") } } return }