package middleware import ( "crypto/rand" "crypto/sha256" "encoding/hex" "fmt" "net/http" "server/internal/config" "server/internal/pkg/errs" "server/internal/pkg/httputil" "strings" "time" "github.com/golang-jwt/jwt/v5" ) type JWTMiddleware struct { cfg *config.JWTConfig } type Claims struct { UserID int32 `json:"user_id"` TokenVersion int32 `json:"token_version"` jwt.RegisteredClaims } func NewJWTMiddleware(cfg *config.Config) (*JWTMiddleware, error) { jwtCfg := cfg.JWTConfig // 不允许空密钥 if jwtCfg.Secret == "" { return nil, fmt.Errorf("jwt.secret 未配置:请在 %s.yaml 中设置 jwt.secret", config.GetEnv()) } m := &JWTMiddleware{cfg: &jwtCfg} // 启动时校验签名算法配置 if _, err := m.signingMethod(); err != nil { return nil, err } return m, nil } // signingMethod 返回配置的签名算法(仅支持 HMAC 家族) func (m *JWTMiddleware) signingMethod() (jwt.SigningMethod, error) { switch strings.ToUpper(m.cfg.SigningMethod) { case "", "HS256": return jwt.SigningMethodHS256, nil case "HS384": return jwt.SigningMethodHS384, nil case "HS512": return jwt.SigningMethodHS512, nil default: return nil, fmt.Errorf("不支持的 jwt.signing_method: %q(仅支持 HS256/HS384/HS512)", m.cfg.SigningMethod) } } // ParseToken 解析accessToken func (m *JWTMiddleware) ParseToken(tokenStr string) (*Claims, error) { token, err := jwt.ParseWithClaims( tokenStr, &Claims{}, func(token *jwt.Token) (any, error) { return []byte(m.cfg.Secret), nil }, ) if err != nil || !token.Valid { return nil, errs.ErrInvalidToken } claims, ok := token.Claims.(*Claims) if !ok { return nil, errs.ErrInvalidTokenClaims } if claims.UserID == 0 { return nil, errs.ErrInvalidTokenClaims } return claims, nil } func (m *JWTMiddleware) Middleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { auth := r.Header.Get("Authorization") if auth == "" { httputil.Fail(w, errs.ErrUnauthenticated) return } parts := strings.SplitN(auth, " ", 2) if len(parts) != 2 || parts[0] != "Bearer" { httputil.Fail(w, errs.ErrUnauthenticated) return } tokenStr := parts[1] claims, err := m.ParseToken(tokenStr) if err != nil { httputil.Fail(w, errs.ErrInvalidToken) return } userCtx := GetUserContext(r.Context()) userCtx.UserID = claims.UserID userCtx.TokenVersion = claims.TokenVersion next.ServeHTTP(w, r) }) } func (m *JWTMiddleware) GenerateAccessToken(userID int32, tokenVersion int32) (string, time.Time, error) { method, err := m.signingMethod() if err != nil { return "", time.Time{}, err } now := time.Now() expiresAt := now.Add(m.cfg.Expire) accessClaims := Claims{ UserID: userID, TokenVersion: tokenVersion, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(expiresAt), IssuedAt: jwt.NewNumericDate(now), }, } // 使用配置的签名算法,与 ParseToken 的校验保持一致 accessToken := jwt.NewWithClaims(method, accessClaims) token, err := accessToken.SignedString([]byte(m.cfg.Secret)) if err != nil { return "", time.Time{}, err } return token, expiresAt, nil } func (m *JWTMiddleware) GenerateRefreshToken() (string, error) { b := make([]byte, 32) _, err := rand.Read(b) if err != nil { return "", err } return hex.EncodeToString(b), nil } // HashRefreshToken 哈希token 存储至redis func (m *JWTMiddleware) HashRefreshToken(token string) string { hash := sha256.Sum256([]byte(token)) return hex.EncodeToString(hash[:]) }