Files
CamTalk/backend/internal/auth/jwt.go
hhs 87c54b80c0 feat: JWT Claims 增加 TokenType 字段,ValidateAccess/ValidateRefresh 区分校验类型
- Claims 新增 TokenType 字段("access" / "refresh")
- GeneratePair 为 access/refresh token 分别设置 token_type
- ValidateAccess 校验后检查 token_type == "access"
- ValidateRefresh 校验后检查 token_type == "refresh"
- 增加 token 类型交叉校验测试
2026-06-20 14:50:19 +08:00

135 lines
3.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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")
)
// 令牌类型常量。
const (
TokenTypeAccess = "access"
TokenTypeRefresh = "refresh"
)
// Claims JWT 声明。
type Claims struct {
UserID string `json:"user_id"`
Username string `json:"username"`
TokenType string `json:"token_type"`
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,
TokenType: TokenTypeAccess,
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,
TokenType: TokenTypeRefresh,
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) {
claims, err := tm.validate(tokenStr)
if err != nil {
return nil, err
}
if claims.TokenType != TokenTypeAccess {
return nil, ErrInvalidToken
}
return claims, nil
}
// ValidateRefresh 校验 refresh token 并返回 Claims。
func (tm *TokenManager) ValidateRefresh(tokenStr string) (*Claims, error) {
claims, err := tm.validate(tokenStr)
if err != nil {
return nil, err
}
if claims.TokenType != TokenTypeRefresh {
return nil, ErrInvalidToken
}
return claims, nil
}
// 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[:])
}