topfans/backend/pkg/jwt/jwt.go
zerosaturation 293c7b14ae fix(security): JWT key governance — MustInit fail-fast + atomic.Value
- 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>
2026-07-23 18:50:12 +08:00

161 lines
4.6 KiB
Go
Raw 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 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())
}