- pkg/jwt: 删 public SetSecret; 加 MustInit(secret string) 启动时强制注入, 缺/为空/等于弱默认值时返回 error(运行期不可再改); 密钥用 atomic.Value 存 []byte,所有读走 mustSecret() 原子 Load,消除 SetSecret/ParseToken 并发 data race(go test -race 零告警)。 - gateway main + auth_provider: 启动时 MustInit 读 JWT_SECRET env,失败 fatal。 - scripts/loadgen/seed/tokens: 同步 MustInit。 - .env.example: JWT_SECRET 改为 ≥32 字节 base64 示例(原为空,被 MustInit 立即拒);注释提示生产 MUST replace。 - 测试: 4 个 MustInit 行为 + 1 个 50-goroutine race 覆盖。 - 行为变更: 任何 .env 缺 JWT_SECRET 或用占位 secret 的服务,启动会 panic (这是 fail-fast 期望行为);其余 4 个 .env 文件占位由 ops 单独轮换。 Co-Authored-By: Claude <noreply@anthropic.com>
161 lines
4.6 KiB
Go
161 lines
4.6 KiB
Go
package jwt
|
||
|
||
import (
|
||
"errors"
|
||
"fmt"
|
||
"sync/atomic"
|
||
"time"
|
||
|
||
"github.com/golang-jwt/jwt/v5"
|
||
)
|
||
|
||
const (
|
||
// TokenExpiration Token过期时间(7天)
|
||
TokenExpiration = 7 * 24 * time.Hour
|
||
// MinSecretLen 强制要求 secret 至少 32 字节(HS256 推荐阈值,等价 256 bit)
|
||
MinSecretLen = 32
|
||
// DefaultSecretValue 已知的弱默认值,启动时若仍用此值则 fail-fast。
|
||
// 保留在源码是为了让 MustInit 主动拒绝"忘了改默认"的服务部署。
|
||
DefaultSecretValue = "your-secret-key-change-in-production"
|
||
)
|
||
|
||
var (
|
||
// ErrSecretEmpty secret 为空
|
||
ErrSecretEmpty = errors.New("jwt: secret is empty")
|
||
// ErrSecretTooShort secret 短于 MinSecretLen
|
||
ErrSecretTooShort = errors.New("jwt: secret shorter than 32 bytes")
|
||
// ErrSecretIsDefault secret 等于已知的弱默认值(开发占位符)
|
||
ErrSecretIsDefault = errors.New("jwt: secret is the well-known default; refuse to run")
|
||
// ErrAlreadyInit MustInit 已被调用过,运行期禁止再次初始化
|
||
ErrAlreadyInit = errors.New("jwt: already initialized")
|
||
|
||
// jwtSecret 用 atomic.Value 存 []byte,所有 reader 走 mustSecret() 的 Load。
|
||
// 旧实现是包级可变 []byte + 无锁 SetSecret,go test -race 必报。
|
||
jwtSecret atomic.Value // []byte
|
||
// initDone 标记 MustInit 是否已成功执行;运行期再次调用会被拒。
|
||
initDone atomic.Bool
|
||
)
|
||
|
||
// MustInit 启动时调用一次,把 secret 注入全局;任何校验失败返回 error
|
||
// (由 main.go Fatal 终止进程)。运行期禁止再设。
|
||
//
|
||
// 校验规则:
|
||
// - 非空
|
||
// - 长度 ≥ 32 字节
|
||
// - 不等于代码里硬编码的弱默认值
|
||
//
|
||
// 满足则原子写入 jwtSecret,后续 Generate/Parse 走 atomic.Load,无 race。
|
||
func MustInit(secret string) error {
|
||
if initDone.Load() {
|
||
return ErrAlreadyInit
|
||
}
|
||
if secret == "" {
|
||
return ErrSecretEmpty
|
||
}
|
||
if len(secret) < MinSecretLen {
|
||
return ErrSecretTooShort
|
||
}
|
||
if secret == DefaultSecretValue {
|
||
return ErrSecretIsDefault
|
||
}
|
||
jwtSecret.Store([]byte(secret))
|
||
initDone.Store(true)
|
||
return nil
|
||
}
|
||
|
||
// setSecretInternal 仅供同 package 的测试(jwt_test.go / race_test.go)使用,
|
||
// 绕过 init-once 限制以便多个 test 各自重设;运行期禁止调用。
|
||
func setSecretInternal(s []byte) {
|
||
jwtSecret.Store(s)
|
||
initDone.Store(true)
|
||
}
|
||
|
||
// mustSecret 从 atomic.Value 读密钥;若未初始化则 panic —— 启动时漏调
|
||
// MustInit 的服务在第一个 Generate/Parse 调用处立即崩,不会静默用弱 key。
|
||
func mustSecret() []byte {
|
||
v := jwtSecret.Load()
|
||
if v == nil {
|
||
panic("jwt: MustInit not called")
|
||
}
|
||
return v.([]byte)
|
||
}
|
||
|
||
// Claims JWT Claims结构
|
||
type Claims struct {
|
||
UserID int64 `json:"user_id"`
|
||
StarID int64 `json:"star_id"`
|
||
UpdatedAt int64 `json:"updated_at"`
|
||
jwt.RegisteredClaims
|
||
}
|
||
|
||
// GenerateToken 生成JWT Token
|
||
func GenerateToken(userID, starID int64, updatedAt int64) (string, error) {
|
||
now := time.Now()
|
||
expiresAt := now.Add(TokenExpiration)
|
||
|
||
claims := Claims{
|
||
UserID: userID,
|
||
StarID: starID,
|
||
UpdatedAt: updatedAt,
|
||
RegisteredClaims: jwt.RegisteredClaims{
|
||
IssuedAt: jwt.NewNumericDate(now),
|
||
ExpiresAt: jwt.NewNumericDate(expiresAt),
|
||
},
|
||
}
|
||
|
||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||
tokenString, err := token.SignedString(mustSecret())
|
||
if err != nil {
|
||
return "", fmt.Errorf("failed to sign token: %w", err)
|
||
}
|
||
|
||
return tokenString, nil
|
||
}
|
||
|
||
// ParseToken 解析JWT Token(不验证过期时间,用于刷新Token场景)
|
||
func ParseToken(tokenString string) (*Claims, error) {
|
||
token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (interface{}, error) {
|
||
// 验证签名算法
|
||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
||
}
|
||
return mustSecret(), nil
|
||
})
|
||
|
||
if err != nil {
|
||
return nil, fmt.Errorf("failed to parse token: %w", err)
|
||
}
|
||
|
||
claims, ok := token.Claims.(*Claims)
|
||
if !ok || !token.Valid {
|
||
return nil, errors.New("invalid token claims")
|
||
}
|
||
|
||
return claims, nil
|
||
}
|
||
|
||
// ValidateToken 验证Token(检查签名和过期时间)
|
||
func ValidateToken(tokenString string) (*Claims, error) {
|
||
claims, err := ParseToken(tokenString)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
// 验证过期时间
|
||
if claims.ExpiresAt != nil && claims.ExpiresAt.Time.Before(time.Now()) {
|
||
return nil, errors.New("token expired")
|
||
}
|
||
|
||
return claims, nil
|
||
}
|
||
|
||
// GetExpiresAt 获取Token过期时间戳(毫秒)
|
||
func GetExpiresAt() int64 {
|
||
return time.Now().Add(TokenExpiration).UnixMilli()
|
||
}
|
||
|
||
// GetExpiresIn 获取Token过期时间(秒)
|
||
func GetExpiresIn() int64 {
|
||
return int64(TokenExpiration.Seconds())
|
||
}
|