refactor(idgen): centralize negative ID handling via panic

Restore NextUint64ID() to uint64-only API so callers need no error
checks. Retry logic stays in idgen; after 3 negative values it panics
instead of returning 0, preventing silent NULL primary keys.
This commit is contained in:
ryan
2026-06-17 13:14:17 +08:00
parent 92dc6a59c0
commit 3a05e7f09c
11 changed files with 17 additions and 68 deletions
+1 -7
View File
@@ -386,14 +386,8 @@ func CreateUser(c *gin.Context) {
return return
} }
nextID, err := idgen.NextUint64ID()
if err != nil {
c.JSON(http.StatusInternalServerError, response.Err(createUserFailed))
return
}
newUser := model.User{ newUser := model.User{
ID: nextID, ID: idgen.NextUint64ID(),
Username: req.Username, Username: req.Username,
Nickname: req.Nickname, Nickname: req.Nickname,
Email: req.Email, Email: req.Email,
+1 -8
View File
@@ -15,7 +15,6 @@ import (
"github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/util" "github.com/Rain-kl/Wavelet/internal/util"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -70,14 +69,8 @@ func RiskControlMiddleware() gin.HandlerFunc {
status = maxHTTPStatus 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{ logItem := &UserAccessLog{
ID: logID, ID: idgen.NextUint64ID(),
UserID: userObj.ID, // 直接从 Context 获取已登录用户ID,避免数据库查询 UserID: userObj.ID, // 直接从 Context 获取已登录用户ID,避免数据库查询
Path: c.Request.URL.Path, Path: c.Request.URL.Path,
Method: c.Request.Method, Method: c.Request.Method,
+2 -10
View File
@@ -140,11 +140,7 @@ func UploadFile(c *gin.Context) {
return return
} }
id, err := idgen.NextUint64ID() id := 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) subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
// 8. 写入当前活动存储驱动。 // 8. 写入当前活动存储驱动。
@@ -376,11 +372,7 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User,
return true, nil return true, nil
} }
id, err := idgen.NextUint64ID() id := idgen.NextUint64ID()
if err != nil {
c.JSON(http.StatusOK, response.Err(ErrSaveUploadRecordFailed))
return true, err
}
newUpload := model.Upload{ newUpload := model.Upload{
ID: id, ID: id,
UserID: currUser.ID, UserID: currUser.ID,
+1 -7
View File
@@ -228,14 +228,8 @@ func Register(c *gin.Context) {
return return
} }
nextID, err := idgen.NextUint64ID()
if err != nil {
c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
return
}
user := model.User{ user := model.User{
ID: nextID, ID: idgen.NextUint64ID(),
Username: req.Username, Username: req.Username,
Nickname: req.Nickname, Nickname: req.Nickname,
Email: req.Email, Email: req.Email,
+4 -8
View File
@@ -6,7 +6,6 @@
package idgen package idgen
import ( import (
"errors"
"fmt" "fmt"
"log" "log"
@@ -19,9 +18,6 @@ const epoch int64 = 1764547200000
const maxNegativeIDRetries = 3 const maxNegativeIDRetries = 3
// ErrNegativeSnowflakeID 表示 Snowflake 在重试后仍生成负值 ID。
var ErrNegativeSnowflakeID = errors.New("snowflake generated negative ID")
var node *snowflake.Node var node *snowflake.Node
func init() { func init() {
@@ -37,14 +33,14 @@ func init() {
} }
// NextUint64ID 生成下一个分布式唯一 ID。 // NextUint64ID 生成下一个分布式唯一 ID。
// 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则返回错误。 // 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则 panic。
func NextUint64ID() (uint64, error) { func NextUint64ID() uint64 {
for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ { for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ {
id := node.Generate().Int64() id := node.Generate().Int64()
if id >= 0 { if id >= 0 {
return uint64(id), nil return uint64(id)
} }
log.Printf("[Snowflake] generated negative ID: %d (attempt %d/%d)", id, attempt, maxNegativeIDRetries) log.Printf("[Snowflake] generated negative ID: %d (attempt %d/%d)", id, attempt, maxNegativeIDRetries)
} }
return 0, fmt.Errorf("%w: failed after %d attempts", ErrNegativeSnowflakeID, maxNegativeIDRetries) panic(fmt.Sprintf("[Snowflake] generated negative ID after %d attempts", maxNegativeIDRetries))
} }
+1 -3
View File
@@ -7,11 +7,9 @@ import (
"testing" "testing"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
) )
func TestNextUint64ID(t *testing.T) { func TestNextUint64ID(t *testing.T) {
id, err := NextUint64ID() id := NextUint64ID()
require.NoError(t, err)
assert.NotZero(t, id) assert.NotZero(t, id)
} }
-1
View File
@@ -5,7 +5,6 @@ package model
const ( const (
errRegistrationDisabled = "注册已关闭" errRegistrationDisabled = "注册已关闭"
errInvalidUserID = "用户 ID 生成失败"
errDatabaseNotInitialized = "database not initialized" errDatabaseNotInitialized = "database not initialized"
errUsernameExists = "用户名已存在" errUsernameExists = "用户名已存在"
errEmailAlreadyBound = "该邮箱已被其他账号绑定" errEmailAlreadyBound = "该邮箱已被其他账号绑定"
+1 -5
View File
@@ -60,11 +60,7 @@ func (TaskExecution) TableName() string {
// CreateTaskExecution 创建任务执行记录 // CreateTaskExecution 创建任务执行记录
func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error { func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error {
id, err := idgen.NextUint64ID() execution.ID = idgen.NextUint64ID()
if err != nil {
return err
}
execution.ID = id
return db.DB(ctx).Create(execution).Error return db.DB(ctx).Create(execution).Error
} }
+1 -5
View File
@@ -134,11 +134,7 @@ func (u *User) assignIDIfMissing() error {
if u.ID != 0 { if u.ID != 0 {
return nil return nil
} }
id, err := idgen.NextUint64ID() u.ID = idgen.NextUint64ID()
if err != nil {
return errors.New(errInvalidUserID)
}
u.ID = id
return nil return nil
} }
+3 -10
View File
@@ -115,10 +115,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
} }
// 生成唯一的 TaskID // 生成唯一的 TaskID
taskID, err := generateTaskID(taskType, triggeredBy) taskID := generateTaskID(taskType, triggeredBy)
if err != nil {
return "", err
}
// 创建任务执行记录 // 创建任务执行记录
execution := &model.TaskExecution{ execution := &model.TaskExecution{
@@ -456,12 +453,8 @@ func handleSuccessfulTask(ctx context.Context, execution *model.TaskExecution, t
} }
// generateTaskID 生成任务 ID // generateTaskID 生成任务 ID
func generateTaskID(taskType string, triggeredBy string) (string, error) { func generateTaskID(taskType string, triggeredBy string) string {
uniqueID, err := idgen.NextUint64ID() return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, idgen.NextUint64ID())
if err != nil {
return "", err
}
return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, uniqueID), nil
} }
// generateRetryTaskID 生成重试任务 ID // generateRetryTaskID 生成重试任务 ID
+2 -4
View File
@@ -387,10 +387,8 @@ func TestRetryTaskNonExistent(t *testing.T) {
} }
func TestGenerateTaskID(t *testing.T) { func TestGenerateTaskID(t *testing.T) {
id1, err := generateTaskID("test_type", "manual") id1 := generateTaskID("test_type", "manual")
require.NoError(t, err) id2 := generateTaskID("test_type", "manual")
id2, err := generateTaskID("test_type", "manual")
require.NoError(t, err)
// 两个 ID 应不同(包含 Snowflake ID) // 两个 ID 应不同(包含 Snowflake ID)
assert.NotEqual(t, id1, id2) assert.NotEqual(t, id1, id2)