package middleware import ( "context" "crypto/rand" "crypto/sha256" "encoding/hex" "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 } 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() (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[:]) }