package middlewares import ( "crypto/rand" "encoding/hex" "fmt" "net/http" "strings" "time" "skeleton/database" "github.com/gin-gonic/gin" "github.com/golang-jwt/jwt/v5" "go.uber.org/zap" ) const ( TokenTypeAccess = "access" TokenTypeRefresh = "refresh" ) var ( jwtSecret []byte jwtIssuer string accessExpire time.Duration refreshExpire time.Duration ) type JWTClaims struct { UserID int `json:"user_id"` Username string `json:"username"` TokenType string `json:"token_type"` jwt.RegisteredClaims } type TokenPair struct { AccessToken string `json:"access_token"` RefreshToken string `json:"refresh_token"` TokenType string `json:"token_type" example:"Bearer"` ExpiresIn int64 `json:"expires_in" example:"900"` } func InitJWT(secret string, accessMinutes, refreshHours int, issuer string) { jwtSecret = []byte(secret) accessExpire = time.Duration(accessMinutes) * time.Minute refreshExpire = time.Duration(refreshHours) * time.Hour jwtIssuer = issuer } func GenerateToken(userID int, username string) (string, error) { return generateToken(userID, username, TokenTypeAccess, accessExpire) } func GenerateTokenPair(userID int, username string) (TokenPair, error) { access, err := generateToken(userID, username, TokenTypeAccess, accessExpire) if err != nil { return TokenPair{}, err } refresh, err := generateToken(userID, username, TokenTypeRefresh, refreshExpire) if err != nil { return TokenPair{}, err } return TokenPair{ AccessToken: access, RefreshToken: refresh, TokenType: "Bearer", ExpiresIn: int64(accessExpire.Seconds()), }, nil } func generateToken(userID int, username, tokenType string, duration time.Duration) (string, error) { now := time.Now() claims := JWTClaims{ UserID: userID, Username: username, TokenType: tokenType, RegisteredClaims: jwt.RegisteredClaims{ ID: randomTokenID(), ExpiresAt: jwt.NewNumericDate(now.Add(duration)), IssuedAt: jwt.NewNumericDate(now), NotBefore: jwt.NewNumericDate(now), Issuer: jwtIssuer, }, } return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(jwtSecret) } func randomTokenID() string { value := make([]byte, 16) if _, err := rand.Read(value); err != nil { return fmt.Sprintf("%d", time.Now().UnixNano()) } return hex.EncodeToString(value) } func ParseToken(tokenString string) (*JWTClaims, error) { token, err := jwt.ParseWithClaims(tokenString, &JWTClaims{}, func(*jwt.Token) (interface{}, error) { return jwtSecret, nil }, jwt.WithValidMethods([]string{jwt.SigningMethodHS256.Alg()}), jwt.WithIssuer(jwtIssuer)) if err != nil { return nil, err } claims, ok := token.Claims.(*JWTClaims) if !ok || !token.Valid { return nil, jwt.ErrTokenInvalidClaims } return claims, nil } func JWTAuth() gin.HandlerFunc { return func(c *gin.Context) { tokenString, err := BearerToken(c.GetHeader("Authorization")) if err != nil { abortUnauthorized(c, err.Error()) return } claims, err := ParseToken(tokenString) if err != nil || claims.TokenType != TokenTypeAccess { abortUnauthorized(c, "token 无效或已过期") return } revoked, err := database.IsTokenRevoked(c.Request.Context(), claims.ID) if err != nil || revoked { abortUnauthorized(c, "token 已撤销") return } c.Set("user_id", claims.UserID) c.Set("username", claims.Username) c.Set("jwt_claims", claims) c.Next() } } func BearerToken(header string) (string, error) { parts := strings.Fields(header) if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") { return "", fmt.Errorf("缺少或格式错误的认证 token") } return parts[1], nil } func abortUnauthorized(c *gin.Context, message string) { Logger.Warn("JWT认证失败", zap.String("message", message), zap.String("path", c.Request.URL.Path)) c.AbortWithStatusJSON(http.StatusOK, gin.H{"code": 401, "message": message, "data": nil}) } func GetCurrentUser(c *gin.Context) (int, string, bool) { userID, userExists := c.Get("user_id") username, nameExists := c.Get("username") if !userExists || !nameExists { return 0, "", false } return userID.(int), username.(string), true } func GetCurrentClaims(c *gin.Context) (*JWTClaims, bool) { value, ok := c.Get("jwt_claims") if !ok { return nil, false } claims, ok := value.(*JWTClaims) return claims, ok }