- Claims 新增 TokenType 字段("access" / "refresh") - GeneratePair 为 access/refresh token 分别设置 token_type - ValidateAccess 校验后检查 token_type == "access" - ValidateRefresh 校验后检查 token_type == "refresh" - 增加 token 类型交叉校验测试
135 lines
3.3 KiB
Go
135 lines
3.3 KiB
Go
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[:])
|
||
}
|