package middleware import ( "context" "errors" "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 contextKey string const UserContextKey contextKey = "user" const RefreshTokenType = "refresh" type Claims struct { UserID int32 `json:"user_id"` jwt.RegisteredClaims } type RefreshClaims struct { UserID int32 `json:"user_id"` Type string `json:"type"` jwt.RegisteredClaims } func NewJWTMiddleware(cfg *config.Config) *JWTMiddleware { jwtCfg := cfg.JWTConfig return &JWTMiddleware{cfg: &jwtCfg} } func GetClaims(ctx context.Context) (*Claims, bool) { claims, ok := ctx.Value(UserContextKey).(*Claims) return claims, ok } // ParseToken 解析accessToken func (m *JWTMiddleware) ParseToken(tokenStr string) (*Claims, error) { token, err := jwt.ParseWithClaims( tokenStr, &Claims{}, func(token *jwt.Token) (interface{}, 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 } // ParseRefreshToken 解析RefreshToken func (m *JWTMiddleware) ParseRefreshToken(tokenStr string) (*RefreshClaims, error) { token, err := jwt.ParseWithClaims( tokenStr, &RefreshClaims{}, func(token *jwt.Token) (interface{}, error) { return []byte(m.cfg.Secret), nil }, ) if err != nil { // Token 已过期 if errors.Is(err, jwt.ErrTokenExpired) { return nil, errs.ErrExpiredRefreshToken } // 签名错误、格式错误、非法 Token return nil, errs.ErrInvalidRefreshToken } if !token.Valid { return nil, errs.ErrInvalidRefreshToken } claims, ok := token.Claims.(*RefreshClaims) if !ok { return nil, errs.ErrInvalidRefreshToken } if claims.UserID == 0 { return nil, errs.ErrInvalidRefreshToken } if claims.Type != RefreshTokenType { return nil, errs.ErrInvalidRefreshToken } 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 } ctx := context.WithValue(r.Context(), UserContextKey, claims) next.ServeHTTP(w, r.WithContext(ctx)) }) } func (m *JWTMiddleware) GenerateAccessToken(userID int32) (string, time.Time, error) { now := time.Now() expiresAt := now.Add(m.cfg.Expire) accessClaims := Claims{ UserID: userID, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(expiresAt), IssuedAt: jwt.NewNumericDate(now), }, } accessToken := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims) token, err := accessToken.SignedString([]byte(m.cfg.Secret)) if err != nil { return "", time.Time{}, err } return token, expiresAt, nil } func (m *JWTMiddleware) GenerateRefreshToken(userID int32) (string, time.Time, error) { now := time.Now() expiresAt := now.Add(m.cfg.RefreshExpire) refreshClaims := RefreshClaims{ UserID: userID, Type: RefreshTokenType, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(expiresAt), IssuedAt: jwt.NewNumericDate(now), }, } refreshToken := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims) token, err := refreshToken.SignedString([]byte(m.cfg.Secret)) if err != nil { return "", time.Time{}, err } return token, expiresAt, nil }