diff --git a/backend/go.mod b/backend/go.mod index 069496b..d14bd93 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -28,6 +28,7 @@ require ( github.com/go-playground/validator/v10 v10.20.0 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/goccy/go-json v0.10.2 // indirect + github.com/golang-jwt/jwt/v5 v5.3.1 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect diff --git a/backend/go.sum b/backend/go.sum index d4c2b90..38b2a4b 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -37,6 +37,8 @@ github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9L github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU= github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI= github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= diff --git a/backend/internal/auth/jwt.go b/backend/internal/auth/jwt.go new file mode 100644 index 0000000..1da1ac0 --- /dev/null +++ b/backend/internal/auth/jwt.go @@ -0,0 +1,111 @@ +package auth + +import ( + "crypto/sha256" + "encoding/hex" + "errors" + "time" + + "github.com/golang-jwt/jwt/v5" + "github.com/google/uuid" +) + +// 自定义错误。 +var ( + ErrInvalidToken = errors.New("invalid or expired token") +) + +// Claims JWT 声明。 +type Claims struct { + UserID string `json:"user_id"` + Username string `json:"username"` + jwt.RegisteredClaims +} + +// TokenManager JWT 令牌管理器。 +type TokenManager struct { + secret []byte + accessTTL time.Duration + refreshTTL time.Duration +} + +// NewTokenManager 创建 TokenManager。 +// secret: JWT 签名密钥;accessTTL/refreshTTL: 令牌有效期。 +func NewTokenManager(secret string, accessTTL, refreshTTL time.Duration) *TokenManager { + return &TokenManager{ + secret: []byte(secret), + accessTTL: accessTTL, + refreshTTL: refreshTTL, + } +} + +// GeneratePair 生成 access + refresh 令牌对。 +func (tm *TokenManager) GeneratePair(userID, username string) (access, refresh string, err error) { + now := time.Now() + + // access token + accessClaims := &Claims{ + UserID: userID, + Username: username, + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(now.Add(tm.accessTTL)), + IssuedAt: jwt.NewNumericDate(now), + Issuer: "camtalk", + }, + } + accessTkn := jwt.NewWithClaims(jwt.SigningMethodHS256, accessClaims) + access, err = accessTkn.SignedString(tm.secret) + if err != nil { + return "", "", err + } + + // refresh token(含唯一 token_id 用于 DB 关联) + tokenID := uuid.New().String() + refreshClaims := &Claims{ + UserID: userID, + Username: username, + RegisteredClaims: jwt.RegisteredClaims{ + ID: tokenID, + ExpiresAt: jwt.NewNumericDate(now.Add(tm.refreshTTL)), + IssuedAt: jwt.NewNumericDate(now), + Issuer: "camtalk", + }, + } + refreshTkn := jwt.NewWithClaims(jwt.SigningMethodHS256, refreshClaims) + refresh, err = refreshTkn.SignedString(tm.secret) + return +} + +// ValidateAccess 校验 access token 并返回 Claims。 +func (tm *TokenManager) ValidateAccess(tokenStr string) (*Claims, error) { + return tm.validate(tokenStr) +} + +// ValidateRefresh 校验 refresh token 并返回 Claims。 +func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) { + return tm.validate(tokenStr) +} + +// validate 解析并校验 JWT。 +func (tm *TokenManager) validate(tokenStr string) (*Claims, error) { + token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(t *jwt.Token) (interface{}, error) { + if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, ErrInvalidToken + } + return tm.secret, nil + }) + if err != nil { + return nil, ErrInvalidToken + } + claims, ok := token.Claims.(*Claims) + if !ok || !token.Valid { + return nil, ErrInvalidToken + } + return claims, nil +} + +// HashToken 对 token 做 SHA256 哈希,用于 DB 存储。 +func HashToken(token string) string { + h := sha256.Sum256([]byte(token)) + return hex.EncodeToString(h[:]) +}