From 92dc6a59c09a6b6afb7fc69a34885f3a9455cf66 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 17 Jun 2026 13:08:00 +0800 Subject: [PATCH] fix(idgen): retry snowflake generation and error on negative ID NextUint64ID now retries up to 3 times when Int64() is negative, then returns an error instead of 0 or Fatalf. All call sites propagate the error to avoid GORM omitting zero-value primary keys. --- internal/apps/admin/user/routers.go | 25 ++++++++++--- internal/apps/risk_control/middleware.go | 16 ++++++-- internal/apps/upload/routers.go | 47 +++++++++++++++--------- internal/apps/user/routers.go | 8 +++- internal/db/idgen/snowflake.go | 25 +++++++++---- internal/db/idgen/snowflake_test.go | 17 +++++++++ internal/model/errs.go | 1 + internal/model/task_execution.go | 6 ++- internal/model/users.go | 24 +++++++++++- internal/task/executor.go | 14 +++++-- internal/task/executor_test.go | 6 ++- 11 files changed, 145 insertions(+), 44 deletions(-) create mode 100644 internal/db/idgen/snowflake_test.go diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go index 00b6b096..f2ee3f43 100644 --- a/internal/apps/admin/user/routers.go +++ b/internal/apps/admin/user/routers.go @@ -4,7 +4,8 @@ package user -import ("net/http" +import ( + "net/http" "strconv" "strings" "time" @@ -17,7 +18,8 @@ import ("net/http" "github.com/gin-gonic/gin" "gorm.io/gorm" - "github.com/Rain-kl/Wavelet/internal/common/response") + "github.com/Rain-kl/Wavelet/internal/common/response" +) // minPasswordLength 密码最小长度 const minPasswordLength = 8 @@ -31,7 +33,7 @@ type listUsersRequest struct { } type user struct { - ID uint64 `json:"id"` + ID uint64 `json:"id,string"` Username string `json:"username"` Nickname string `json:"nickname"` Email string `json:"email"` @@ -104,7 +106,7 @@ func ListUsers(c *gin.Context) { return } - var users []user + var modelUsers []model.User var total int64 query := db.DB(c.Request.Context()).Model(&model.User{}) @@ -131,11 +133,16 @@ func ListUsers(c *gin.Context) { Order("id DESC"). Offset(offset). Limit(req.PageSize). - Find(&users).Error; err != nil { + Find(&modelUsers).Error; err != nil { c.JSON(http.StatusInternalServerError, response.Err(err.Error())) return } + users := make([]user, 0, len(modelUsers)) + for _, modelUser := range modelUsers { + users = append(users, toUser(modelUser)) + } + c.JSON(http.StatusOK, response.OK(listUsersResponse{ Users: users, Total: total, @@ -379,8 +386,14 @@ func CreateUser(c *gin.Context) { return } + nextID, err := idgen.NextUint64ID() + if err != nil { + c.JSON(http.StatusInternalServerError, response.Err(createUserFailed)) + return + } + newUser := model.User{ - ID: idgen.NextUint64ID(), + ID: nextID, Username: req.Username, Nickname: req.Nickname, Email: req.Email, diff --git a/internal/apps/risk_control/middleware.go b/internal/apps/risk_control/middleware.go index de13775d..ead5e526 100644 --- a/internal/apps/risk_control/middleware.go +++ b/internal/apps/risk_control/middleware.go @@ -4,18 +4,20 @@ // Package risk_control 提供风险控制中间件 package risk_control -import ("encoding/json" +import ( + "encoding/json" "net/http" "time" "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/config" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/util" + "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" - - "github.com/Rain-kl/Wavelet/internal/common/response") +) // RiskControlMiddleware 全局日志采集中间件 func RiskControlMiddleware() gin.HandlerFunc { @@ -68,8 +70,14 @@ func RiskControlMiddleware() gin.HandlerFunc { status = maxHTTPStatus } + logID, err := idgen.NextUint64ID() + if err != nil { + logger.ErrorF(c.Request.Context(), "[RiskControl] access log ID generation failed: %v", err) + return + } + logItem := &UserAccessLog{ - ID: idgen.NextUint64ID(), + ID: logID, UserID: userObj.ID, // 直接从 Context 获取已登录用户ID,避免数据库查询 Path: c.Request.URL.Path, Method: c.Request.Method, diff --git a/internal/apps/upload/routers.go b/internal/apps/upload/routers.go index be9209e0..45a86002 100644 --- a/internal/apps/upload/routers.go +++ b/internal/apps/upload/routers.go @@ -117,21 +117,10 @@ func UploadFile(c *gin.Context) { uploadType := c.DefaultPostForm("type", "generic") - accessModeStr := c.PostForm("access_mode") - var accessMode int - if accessModeStr == "" { - if uploadType == defaultPublicUploadType { - accessMode = 1 - } else { - accessMode = 0 - } - } else { - var err error - accessMode, err = strconv.Atoi(accessModeStr) - if err != nil || (accessMode != 0 && accessMode != 1) { - c.JSON(http.StatusOK, response.Err("无效的 access_mode 参数")) - return - } + accessMode, errMsg := resolveUploadAccessMode(c, uploadType) + if errMsg != "" { + c.JSON(http.StatusOK, response.Err(errMsg)) + return } // 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件 @@ -151,7 +140,11 @@ func UploadFile(c *gin.Context) { return } - id := idgen.NextUint64ID() + id, err := idgen.NextUint64ID() + if err != nil { + c.JSON(http.StatusOK, response.Err(ErrSaveUploadRecordFailed)) + return + } subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext) // 8. 写入当前活动存储驱动。 @@ -336,6 +329,22 @@ func BatchDownloadFiles(c *gin.Context) { } } +func resolveUploadAccessMode(c *gin.Context, uploadType string) (int, string) { + accessModeStr := c.PostForm("access_mode") + if accessModeStr == "" { + if uploadType == defaultPublicUploadType { + return 1, "" + } + return 0, "" + } + + accessMode, err := strconv.Atoi(accessModeStr) + if err != nil || (accessMode != 0 && accessMode != 1) { + return 0, "无效的 access_mode 参数" + } + return accessMode, "" +} + // validateUploadExtension 校验文件后缀是否在系统允许的上传扩展名列表中 func validateUploadExtension(ctx context.Context, ext string) string { var sc model.SystemConfig @@ -367,7 +376,11 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User, return true, nil } - id := idgen.NextUint64ID() + id, err := idgen.NextUint64ID() + if err != nil { + c.JSON(http.StatusOK, response.Err(ErrSaveUploadRecordFailed)) + return true, err + } newUpload := model.Upload{ ID: id, UserID: currUser.ID, diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index fad35df2..df078115 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -228,8 +228,14 @@ func Register(c *gin.Context) { return } + nextID, err := idgen.NextUint64ID() + if err != nil { + c.JSON(http.StatusInternalServerError, response.Err(err.Error())) + return + } + user := model.User{ - ID: idgen.NextUint64ID(), + ID: nextID, Username: req.Username, Nickname: req.Nickname, Email: req.Email, diff --git a/internal/db/idgen/snowflake.go b/internal/db/idgen/snowflake.go index 8f67fc87..467aebee 100644 --- a/internal/db/idgen/snowflake.go +++ b/internal/db/idgen/snowflake.go @@ -6,6 +6,8 @@ package idgen import ( + "errors" + "fmt" "log" "github.com/Rain-kl/Wavelet/internal/config" @@ -15,6 +17,11 @@ import ( // 2025-12-01 00:00:00 UTC 的毫秒时间戳 const epoch int64 = 1764547200000 +const maxNegativeIDRetries = 3 + +// ErrNegativeSnowflakeID 表示 Snowflake 在重试后仍生成负值 ID。 +var ErrNegativeSnowflakeID = errors.New("snowflake generated negative ID") + var node *snowflake.Node func init() { @@ -29,11 +36,15 @@ func init() { log.Printf("[Snowflake] initialized with node ID: %d, epoch: 2025-12-01\n", nodeID) } -// NextUint64ID 生成下一个分布式唯一 ID -func NextUint64ID() uint64 { - val := node.Generate().Int64() - if val < 0 { - return 0 +// NextUint64ID 生成下一个分布式唯一 ID。 +// 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则返回错误。 +func NextUint64ID() (uint64, error) { + for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ { + id := node.Generate().Int64() + if id >= 0 { + return uint64(id), nil + } + log.Printf("[Snowflake] generated negative ID: %d (attempt %d/%d)", id, attempt, maxNegativeIDRetries) } - return uint64(val) -} + return 0, fmt.Errorf("%w: failed after %d attempts", ErrNegativeSnowflakeID, maxNegativeIDRetries) +} \ No newline at end of file diff --git a/internal/db/idgen/snowflake_test.go b/internal/db/idgen/snowflake_test.go new file mode 100644 index 00000000..72d74513 --- /dev/null +++ b/internal/db/idgen/snowflake_test.go @@ -0,0 +1,17 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package idgen + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestNextUint64ID(t *testing.T) { + id, err := NextUint64ID() + require.NoError(t, err) + assert.NotZero(t, id) +} \ No newline at end of file diff --git a/internal/model/errs.go b/internal/model/errs.go index 6c2a64ee..f9198372 100644 --- a/internal/model/errs.go +++ b/internal/model/errs.go @@ -5,6 +5,7 @@ package model const ( errRegistrationDisabled = "注册已关闭" + errInvalidUserID = "用户 ID 生成失败" errDatabaseNotInitialized = "database not initialized" errUsernameExists = "用户名已存在" errEmailAlreadyBound = "该邮箱已被其他账号绑定" diff --git a/internal/model/task_execution.go b/internal/model/task_execution.go index 81a1d193..37ddb7e8 100644 --- a/internal/model/task_execution.go +++ b/internal/model/task_execution.go @@ -60,7 +60,11 @@ func (TaskExecution) TableName() string { // CreateTaskExecution 创建任务执行记录 func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error { - execution.ID = idgen.NextUint64ID() + id, err := idgen.NextUint64ID() + if err != nil { + return err + } + execution.ID = id return db.DB(ctx).Create(execution).Error } diff --git a/internal/model/users.go b/internal/model/users.go index aa5b83a2..cb91a279 100644 --- a/internal/model/users.go +++ b/internal/model/users.go @@ -12,6 +12,7 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/common" + "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/pkg/util" "gorm.io/gorm" ) @@ -44,7 +45,7 @@ func (u *OAuthUserInfo) GetID() uint64 { // User 用户表实体 type User struct { - ID uint64 `json:"id" gorm:"primaryKey"` + ID uint64 `json:"id,string" gorm:"primaryKey;not null"` Username string `json:"username" gorm:"size:64;uniqueIndex"` Password string `json:"password,omitempty" gorm:"size:255"` Nickname string `json:"nickname" gorm:"size:255"` @@ -129,6 +130,18 @@ func (u *User) CheckActive() error { return nil } +func (u *User) assignIDIfMissing() error { + if u.ID != 0 { + return nil + } + id, err := idgen.NextUint64ID() + if err != nil { + return errors.New(errInvalidUserID) + } + u.ID = id + return nil +} + // CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验) func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error { enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled) @@ -137,8 +150,9 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser } now := time.Now() + userID := oauthInfo.GetID() newUser := User{ - ID: oauthInfo.GetID(), + ID: userID, Username: oauthInfo.Username, Nickname: oauthInfo.Name, Email: oauthInfo.Email, @@ -147,6 +161,9 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser LastLoginAt: now, IsAdmin: false, } + if err := newUser.assignIDIfMissing(); err != nil { + return err + } if err := tx.Create(&newUser).Error; err != nil { return err } @@ -182,6 +199,9 @@ func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error { } } + if err := u.assignIDIfMissing(); err != nil { + return err + } if err := tx.Create(u).Error; err != nil { return err } diff --git a/internal/task/executor.go b/internal/task/executor.go index 925ab0a1..6a61e37d 100644 --- a/internal/task/executor.go +++ b/internal/task/executor.go @@ -115,7 +115,10 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere } // 生成唯一的 TaskID - taskID := generateTaskID(taskType, triggeredBy) + taskID, err := generateTaskID(taskType, triggeredBy) + if err != nil { + return "", err + } // 创建任务执行记录 execution := &model.TaskExecution{ @@ -453,9 +456,12 @@ func handleSuccessfulTask(ctx context.Context, execution *model.TaskExecution, t } // generateTaskID 生成任务 ID -func generateTaskID(taskType string, triggeredBy string) string { - uniqueID := idgen.NextUint64ID() - return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, uniqueID) +func generateTaskID(taskType string, triggeredBy string) (string, error) { + uniqueID, err := idgen.NextUint64ID() + if err != nil { + return "", err + } + return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, uniqueID), nil } // generateRetryTaskID 生成重试任务 ID diff --git a/internal/task/executor_test.go b/internal/task/executor_test.go index fc53b770..674ebd1d 100644 --- a/internal/task/executor_test.go +++ b/internal/task/executor_test.go @@ -387,8 +387,10 @@ func TestRetryTaskNonExistent(t *testing.T) { } func TestGenerateTaskID(t *testing.T) { - id1 := generateTaskID("test_type", "manual") - id2 := generateTaskID("test_type", "manual") + id1, err := generateTaskID("test_type", "manual") + require.NoError(t, err) + id2, err := generateTaskID("test_type", "manual") + require.NoError(t, err) // 两个 ID 应不同(包含 Snowflake ID) assert.NotEqual(t, id1, id2)