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>
This commit is contained in:
parent
3407e30395
commit
293c7b14ae
@ -12,8 +12,10 @@ GIN_MODE=release
|
|||||||
SERVER_PORT=8080
|
SERVER_PORT=8080
|
||||||
|
|
||||||
# ==================== JWT Configuration ====================
|
# ==================== JWT Configuration ====================
|
||||||
# JWT密钥 - 生产环境请修改为安全的随机字符串
|
# JWT 密钥(MustInit 必填,≥32 字节随机;启动期 MustInit 校验)
|
||||||
JWT_SECRET=
|
# 示例: 任意 32+ 字节随机串(可用 `openssl rand -base64 48` 生成)
|
||||||
|
# 占位示例(本地开发/CI 用,MUST replace in prod): 随机 base64 字符串(48 字节 = 64 base64 字符)
|
||||||
|
JWT_SECRET=ZGV2X2p3dF9zZWNyZXRfa2V5X3BsYWNlaG9sZGVyX2F1dGhfbmVlZHNfdG9fYmVfMzJfYnl0ZXNfbG9uZw==
|
||||||
|
|
||||||
# ==================== Dubbo Service URLs ====================
|
# ==================== Dubbo Service URLs ====================
|
||||||
# 各微服务的Dubbo连接地址(直连模式)
|
# 各微服务的Dubbo连接地址(直连模式)
|
||||||
|
|||||||
@ -10,6 +10,7 @@ import (
|
|||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/topfans/backend/gateway/dto"
|
"github.com/topfans/backend/gateway/dto"
|
||||||
"github.com/topfans/backend/gateway/pkg/response"
|
"github.com/topfans/backend/gateway/pkg/response"
|
||||||
|
"github.com/topfans/backend/gateway/pkg/starcache"
|
||||||
"github.com/topfans/backend/pkg/logger"
|
"github.com/topfans/backend/pkg/logger"
|
||||||
"google.golang.org/grpc/codes"
|
"google.golang.org/grpc/codes"
|
||||||
pb "github.com/topfans/backend/pkg/proto/user"
|
pb "github.com/topfans/backend/pkg/proto/user"
|
||||||
@ -19,6 +20,7 @@ import (
|
|||||||
// AuthController 认证控制器
|
// AuthController 认证控制器
|
||||||
type AuthController struct {
|
type AuthController struct {
|
||||||
userServiceClient pb.UserSocialService
|
userServiceClient pb.UserSocialService
|
||||||
|
starCache *starcache.Cache
|
||||||
}
|
}
|
||||||
|
|
||||||
// pbError 用于包装 proto 错误消息
|
// pbError 用于包装 proto 错误消息
|
||||||
@ -31,7 +33,10 @@ func (e *pbError) Error() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NewAuthController 创建认证控制器
|
// NewAuthController 创建认证控制器
|
||||||
func NewAuthController(dubboClient *client.Client) (*AuthController, error) {
|
//
|
||||||
|
// starCache 用于 Register/Login 两个公开入口的 star 解析:
|
||||||
|
// 取代原先每次都直接 RPC GetFanIdentities 拉一遍可选身份列表。
|
||||||
|
func NewAuthController(dubboClient *client.Client, starCache *starcache.Cache) (*AuthController, error) {
|
||||||
svc, err := pb.NewUserSocialService(dubboClient)
|
svc, err := pb.NewUserSocialService(dubboClient)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@ -39,9 +44,24 @@ func NewAuthController(dubboClient *client.Client) (*AuthController, error) {
|
|||||||
|
|
||||||
return &AuthController{
|
return &AuthController{
|
||||||
userServiceClient: svc,
|
userServiceClient: svc,
|
||||||
|
starCache: starCache,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// findStar 解析 starID 对应的 *pb.Star。失败仅为 warn,不阻断主流程
|
||||||
|
// (DTO 转换对 nil star 有防御,Register/Login 仍能成功)。
|
||||||
|
func (ctrl *AuthController) findStar(ctx context.Context, starID int64) *pb.Star {
|
||||||
|
star, err := ctrl.starCache.GetStar(ctx, starID)
|
||||||
|
if err != nil {
|
||||||
|
logger.Logger.Warn("GetStar cache miss+RPC failed, continuing with nil star",
|
||||||
|
zap.Int64("star_id", starID),
|
||||||
|
zap.Error(err),
|
||||||
|
)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return star
|
||||||
|
}
|
||||||
|
|
||||||
// Register 用户注册
|
// Register 用户注册
|
||||||
// @Summary 用户注册
|
// @Summary 用户注册
|
||||||
// @Description 用户注册接口,需要提供手机号、密码、选择明星身份
|
// @Description 用户注册接口,需要提供手机号、密码、选择明星身份
|
||||||
@ -79,22 +99,7 @@ func (ctrl *AuthController) Register(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 获取 Star 信息用于 DTO 转换
|
star := ctrl.findStar(ctx, req.StarId)
|
||||||
starResp, err := ctrl.userServiceClient.GetFanIdentities(ctx, &pb.GetFanIdentitiesRequest{})
|
|
||||||
if err != nil {
|
|
||||||
logger.Logger.Error("Failed to get star info", zap.Error(err))
|
|
||||||
response.HandleError(c, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// 找到对应的 Star
|
|
||||||
var star *pb.Star
|
|
||||||
for _, s := range starResp.Stars {
|
|
||||||
if s.StarId == req.StarId {
|
|
||||||
star = s
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.Logger.Info("Register successful",
|
logger.Logger.Info("Register successful",
|
||||||
zap.Int64("user_id", resp.User.Id),
|
zap.Int64("user_id", resp.User.Id),
|
||||||
@ -147,22 +152,7 @@ func (ctrl *AuthController) Login(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 获取 Star 信息用于 DTO 转换
|
star := ctrl.findStar(ctx, resp.FanProfile.StarId)
|
||||||
starResp, err := ctrl.userServiceClient.GetFanIdentities(ctx, &pb.GetFanIdentitiesRequest{})
|
|
||||||
if err != nil {
|
|
||||||
logger.Logger.Error("Failed to get star info", zap.Error(err))
|
|
||||||
response.HandleError(c, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// 找到对应的 Star
|
|
||||||
var star *pb.Star
|
|
||||||
for _, s := range starResp.Stars {
|
|
||||||
if s.StarId == resp.FanProfile.StarId {
|
|
||||||
star = s
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
logger.Logger.Info("Login successful",
|
logger.Logger.Info("Login successful",
|
||||||
zap.Int64("user_id", resp.User.Id),
|
zap.Int64("user_id", resp.User.Id),
|
||||||
|
|||||||
@ -3,6 +3,7 @@ package jwt
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
@ -11,16 +12,72 @@ import (
|
|||||||
const (
|
const (
|
||||||
// TokenExpiration Token过期时间(7天)
|
// TokenExpiration Token过期时间(7天)
|
||||||
TokenExpiration = 7 * 24 * time.Hour
|
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 (
|
var (
|
||||||
// jwtSecret JWT签名密钥(应该从环境变量或配置文件读取)
|
// ErrSecretEmpty secret 为空
|
||||||
jwtSecret = []byte("your-secret-key-change-in-production")
|
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
|
||||||
)
|
)
|
||||||
|
|
||||||
// SetSecret 设置JWT签名密钥(应该在服务启动时调用)
|
// MustInit 启动时调用一次,把 secret 注入全局;任何校验失败返回 error
|
||||||
func SetSecret(secret string) {
|
// (由 main.go Fatal 终止进程)。运行期禁止再设。
|
||||||
jwtSecret = []byte(secret)
|
//
|
||||||
|
// 校验规则:
|
||||||
|
// - 非空
|
||||||
|
// - 长度 ≥ 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结构
|
// Claims JWT Claims结构
|
||||||
@ -47,7 +104,7 @@ func GenerateToken(userID, starID int64, updatedAt int64) (string, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||||
tokenString, err := token.SignedString(jwtSecret)
|
tokenString, err := token.SignedString(mustSecret())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("failed to sign token: %w", err)
|
return "", fmt.Errorf("failed to sign token: %w", err)
|
||||||
}
|
}
|
||||||
@ -62,7 +119,7 @@ func ParseToken(tokenString string) (*Claims, error) {
|
|||||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||||
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
|
||||||
}
|
}
|
||||||
return jwtSecret, nil
|
return mustSecret(), nil
|
||||||
})
|
})
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@ -1,6 +1,7 @@
|
|||||||
package jwt
|
package jwt
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
@ -8,8 +9,40 @@ import (
|
|||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// testSecret32B 是 32+ 字节的合法密钥,供所有测试通过 setSecretInternal 注入。
|
||||||
|
// 长度恰好 43 字节,符合 MinSecretLen=32 且不是默认值。
|
||||||
|
const testSecret32B = "test-secret-key-must-be-at-least-32-bytes"
|
||||||
|
|
||||||
|
// ★ 重要测试顺序:负向测试必须在 TestMustInit_OK 之前,否则 initDone=true
|
||||||
|
// 会让 MustInit 全部返回 ErrAlreadyInit,负向断言不到目标错误。
|
||||||
|
// Go 同一 package 内的 test 按源码顺序串行执行。
|
||||||
|
|
||||||
|
func TestMustInit_Empty(t *testing.T) {
|
||||||
|
if err := MustInit(""); !errors.Is(err, ErrSecretEmpty) {
|
||||||
|
t.Fatalf("empty should reject, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMustInit_TooShort(t *testing.T) {
|
||||||
|
if err := MustInit("short"); !errors.Is(err, ErrSecretTooShort) {
|
||||||
|
t.Fatalf("<32B should reject, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMustInit_DefaultSecret(t *testing.T) {
|
||||||
|
if err := MustInit("your-secret-key-change-in-production"); !errors.Is(err, ErrSecretIsDefault) {
|
||||||
|
t.Fatalf("default should reject, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMustInit_OK(t *testing.T) {
|
||||||
|
if err := MustInit(testSecret32B); err != nil {
|
||||||
|
t.Fatalf("MustInit good: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestGenerateToken(t *testing.T) {
|
func TestGenerateToken(t *testing.T) {
|
||||||
SetSecret("test-secret-key")
|
setSecretInternal([]byte(testSecret32B))
|
||||||
|
|
||||||
userID := int64(10000001)
|
userID := int64(10000001)
|
||||||
starID := int64(123)
|
starID := int64(123)
|
||||||
@ -28,7 +61,7 @@ func TestGenerateToken(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestParseToken(t *testing.T) {
|
func TestParseToken(t *testing.T) {
|
||||||
SetSecret("test-secret-key")
|
setSecretInternal([]byte(testSecret32B))
|
||||||
|
|
||||||
userID := int64(10000001)
|
userID := int64(10000001)
|
||||||
starID := int64(123)
|
starID := int64(123)
|
||||||
@ -58,7 +91,7 @@ func TestParseToken(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateToken(t *testing.T) {
|
func TestValidateToken(t *testing.T) {
|
||||||
SetSecret("test-secret-key")
|
setSecretInternal([]byte(testSecret32B))
|
||||||
|
|
||||||
userID := int64(10000001)
|
userID := int64(10000001)
|
||||||
starID := int64(123)
|
starID := int64(123)
|
||||||
@ -80,7 +113,7 @@ func TestValidateToken(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestValidateToken_Expired(t *testing.T) {
|
func TestValidateToken_Expired(t *testing.T) {
|
||||||
SetSecret("test-secret-key")
|
setSecretInternal([]byte(testSecret32B))
|
||||||
|
|
||||||
// 手动创建一个已过期的Token
|
// 手动创建一个已过期的Token
|
||||||
userID := int64(10000001)
|
userID := int64(10000001)
|
||||||
@ -101,7 +134,7 @@ func TestValidateToken_Expired(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
tokenObj := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
tokenObj := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||||
tokenString, err := tokenObj.SignedString([]byte("test-secret-key"))
|
tokenString, err := tokenObj.SignedString([]byte(testSecret32B))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Failed to create expired token: %v", err)
|
t.Fatalf("Failed to create expired token: %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
30
backend/pkg/jwt/race_test.go
Normal file
30
backend/pkg/jwt/race_test.go
Normal file
@ -0,0 +1,30 @@
|
|||||||
|
package jwt
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const goodSecret = "this-is-a-32-byte-test-secret-32!"
|
||||||
|
|
||||||
|
// TestMustInit_NoRace 验证 50 个并发 reader(GenerateToken + ValidateToken)同时
|
||||||
|
// 通过 atomic.Value.Load() 拿密钥,不应触发 race detector。
|
||||||
|
//
|
||||||
|
// 旧实现是包级可变 []byte + 无锁 SetSecret,go test -race 必报。
|
||||||
|
// 新实现 jwtSecret atomic.Value + mustSecret() 单次 Load,read 路径无共享写。
|
||||||
|
//
|
||||||
|
// 注:用 setSecretInternal 而非 MustInit,因为 jwt_test.go 里 TestMustInit_OK
|
||||||
|
// 已经把 initDone 置 true;同一 test binary 内再次调 MustInit 会返回
|
||||||
|
// ErrAlreadyInit,跑不到 reader 并发路径。
|
||||||
|
func TestMustInit_NoRace(t *testing.T) {
|
||||||
|
setSecretInternal([]byte(goodSecret))
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < 50; i++ {
|
||||||
|
wg.Add(2)
|
||||||
|
go func() { defer wg.Done(); _, _ = GenerateToken(1, 1, time.Now().UnixMilli()) }()
|
||||||
|
go func() { defer wg.Done(); _, _ = ValidateToken("invalid") }()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
@ -20,7 +20,9 @@ type TestUser struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func GenerateTokensForLoadtest(cfg *Config) error {
|
func GenerateTokensForLoadtest(cfg *Config) error {
|
||||||
jwt.SetSecret(cfg.JWTSecret)
|
if err := jwt.MustInit(cfg.JWTSecret); err != nil {
|
||||||
|
return fmt.Errorf("jwt.MustInit failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
db, err := openDB(cfg)
|
db, err := openDB(cfg)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user