mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-09 00:56:37 +08:00
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.
This commit is contained in:
@@ -4,7 +4,8 @@
|
|||||||
|
|
||||||
package user
|
package user
|
||||||
|
|
||||||
import ("net/http"
|
import (
|
||||||
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -17,7 +18,8 @@ import ("net/http"
|
|||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
|
)
|
||||||
|
|
||||||
// minPasswordLength 密码最小长度
|
// minPasswordLength 密码最小长度
|
||||||
const minPasswordLength = 8
|
const minPasswordLength = 8
|
||||||
@@ -31,7 +33,7 @@ type listUsersRequest struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type user struct {
|
type user struct {
|
||||||
ID uint64 `json:"id"`
|
ID uint64 `json:"id,string"`
|
||||||
Username string `json:"username"`
|
Username string `json:"username"`
|
||||||
Nickname string `json:"nickname"`
|
Nickname string `json:"nickname"`
|
||||||
Email string `json:"email"`
|
Email string `json:"email"`
|
||||||
@@ -104,7 +106,7 @@ func ListUsers(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
var users []user
|
var modelUsers []model.User
|
||||||
var total int64
|
var total int64
|
||||||
|
|
||||||
query := db.DB(c.Request.Context()).Model(&model.User{})
|
query := db.DB(c.Request.Context()).Model(&model.User{})
|
||||||
@@ -131,11 +133,16 @@ func ListUsers(c *gin.Context) {
|
|||||||
Order("id DESC").
|
Order("id DESC").
|
||||||
Offset(offset).
|
Offset(offset).
|
||||||
Limit(req.PageSize).
|
Limit(req.PageSize).
|
||||||
Find(&users).Error; err != nil {
|
Find(&modelUsers).Error; err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
|
c.JSON(http.StatusInternalServerError, response.Err(err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
users := make([]user, 0, len(modelUsers))
|
||||||
|
for _, modelUser := range modelUsers {
|
||||||
|
users = append(users, toUser(modelUser))
|
||||||
|
}
|
||||||
|
|
||||||
c.JSON(http.StatusOK, response.OK(listUsersResponse{
|
c.JSON(http.StatusOK, response.OK(listUsersResponse{
|
||||||
Users: users,
|
Users: users,
|
||||||
Total: total,
|
Total: total,
|
||||||
@@ -379,8 +386,14 @@ 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: idgen.NextUint64ID(),
|
ID: nextID,
|
||||||
Username: req.Username,
|
Username: req.Username,
|
||||||
Nickname: req.Nickname,
|
Nickname: req.Nickname,
|
||||||
Email: req.Email,
|
Email: req.Email,
|
||||||
|
|||||||
@@ -4,18 +4,20 @@
|
|||||||
// Package risk_control 提供风险控制中间件
|
// Package risk_control 提供风险控制中间件
|
||||||
package risk_control
|
package risk_control
|
||||||
|
|
||||||
import ("encoding/json"
|
import (
|
||||||
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"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/config"
|
||||||
"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"
|
||||||
|
)
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response")
|
|
||||||
|
|
||||||
// RiskControlMiddleware 全局日志采集中间件
|
// RiskControlMiddleware 全局日志采集中间件
|
||||||
func RiskControlMiddleware() gin.HandlerFunc {
|
func RiskControlMiddleware() gin.HandlerFunc {
|
||||||
@@ -68,8 +70,14 @@ 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: idgen.NextUint64ID(),
|
ID: logID,
|
||||||
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,
|
||||||
|
|||||||
@@ -117,21 +117,10 @@ func UploadFile(c *gin.Context) {
|
|||||||
|
|
||||||
uploadType := c.DefaultPostForm("type", "generic")
|
uploadType := c.DefaultPostForm("type", "generic")
|
||||||
|
|
||||||
accessModeStr := c.PostForm("access_mode")
|
accessMode, errMsg := resolveUploadAccessMode(c, uploadType)
|
||||||
var accessMode int
|
if errMsg != "" {
|
||||||
if accessModeStr == "" {
|
c.JSON(http.StatusOK, response.Err(errMsg))
|
||||||
if uploadType == defaultPublicUploadType {
|
return
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件
|
// 6. 秒传匹配校验:校验数据库中是否存在相同 Hash 且大小一致的可用文件
|
||||||
@@ -151,7 +140,11 @@ func UploadFile(c *gin.Context) {
|
|||||||
return
|
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)
|
subPath := fmt.Sprintf("uploads/%s/%d.%s", time.Now().Format("2006/01/02"), id, ext)
|
||||||
|
|
||||||
// 8. 写入当前活动存储驱动。
|
// 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 校验文件后缀是否在系统允许的上传扩展名列表中
|
// validateUploadExtension 校验文件后缀是否在系统允许的上传扩展名列表中
|
||||||
func validateUploadExtension(ctx context.Context, ext string) string {
|
func validateUploadExtension(ctx context.Context, ext string) string {
|
||||||
var sc model.SystemConfig
|
var sc model.SystemConfig
|
||||||
@@ -367,7 +376,11 @@ func tryInstantUpload(ctx context.Context, c *gin.Context, currUser *model.User,
|
|||||||
return true, nil
|
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{
|
newUpload := model.Upload{
|
||||||
ID: id,
|
ID: id,
|
||||||
UserID: currUser.ID,
|
UserID: currUser.ID,
|
||||||
|
|||||||
@@ -228,8 +228,14 @@ 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: idgen.NextUint64ID(),
|
ID: nextID,
|
||||||
Username: req.Username,
|
Username: req.Username,
|
||||||
Nickname: req.Nickname,
|
Nickname: req.Nickname,
|
||||||
Email: req.Email,
|
Email: req.Email,
|
||||||
|
|||||||
@@ -6,6 +6,8 @@
|
|||||||
package idgen
|
package idgen
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
"log"
|
"log"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/config"
|
"github.com/Rain-kl/Wavelet/internal/config"
|
||||||
@@ -15,6 +17,11 @@ import (
|
|||||||
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
// 2025-12-01 00:00:00 UTC 的毫秒时间戳
|
||||||
const epoch int64 = 1764547200000
|
const epoch int64 = 1764547200000
|
||||||
|
|
||||||
|
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() {
|
||||||
@@ -29,11 +36,15 @@ func init() {
|
|||||||
log.Printf("[Snowflake] initialized with node ID: %d, epoch: 2025-12-01\n", nodeID)
|
log.Printf("[Snowflake] initialized with node ID: %d, epoch: 2025-12-01\n", nodeID)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NextUint64ID 生成下一个分布式唯一 ID
|
// NextUint64ID 生成下一个分布式唯一 ID。
|
||||||
func NextUint64ID() uint64 {
|
// 理论上不应出现负值;若出现则最多重试 maxNegativeIDRetries 次,仍失败则返回错误。
|
||||||
val := node.Generate().Int64()
|
func NextUint64ID() (uint64, error) {
|
||||||
if val < 0 {
|
for attempt := 1; attempt <= maxNegativeIDRetries; attempt++ {
|
||||||
return 0
|
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)
|
||||||
}
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -5,6 +5,7 @@ package model
|
|||||||
|
|
||||||
const (
|
const (
|
||||||
errRegistrationDisabled = "注册已关闭"
|
errRegistrationDisabled = "注册已关闭"
|
||||||
|
errInvalidUserID = "用户 ID 生成失败"
|
||||||
errDatabaseNotInitialized = "database not initialized"
|
errDatabaseNotInitialized = "database not initialized"
|
||||||
errUsernameExists = "用户名已存在"
|
errUsernameExists = "用户名已存在"
|
||||||
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
||||||
|
|||||||
@@ -60,7 +60,11 @@ func (TaskExecution) TableName() string {
|
|||||||
|
|
||||||
// CreateTaskExecution 创建任务执行记录
|
// CreateTaskExecution 创建任务执行记录
|
||||||
func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
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
|
return db.DB(ctx).Create(execution).Error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+22
-2
@@ -12,6 +12,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common"
|
"github.com/Rain-kl/Wavelet/internal/common"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/db/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/pkg/util"
|
"github.com/Rain-kl/Wavelet/pkg/util"
|
||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
@@ -44,7 +45,7 @@ func (u *OAuthUserInfo) GetID() uint64 {
|
|||||||
|
|
||||||
// User 用户表实体
|
// User 用户表实体
|
||||||
type User struct {
|
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"`
|
Username string `json:"username" gorm:"size:64;uniqueIndex"`
|
||||||
Password string `json:"password,omitempty" gorm:"size:255"`
|
Password string `json:"password,omitempty" gorm:"size:255"`
|
||||||
Nickname string `json:"nickname" gorm:"size:255"`
|
Nickname string `json:"nickname" gorm:"size:255"`
|
||||||
@@ -129,6 +130,18 @@ func (u *User) CheckActive() error {
|
|||||||
return nil
|
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 自动注册,含底层权限校验)
|
// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验)
|
||||||
func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
|
func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
|
||||||
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
|
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
|
||||||
@@ -137,8 +150,9 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser
|
|||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
|
userID := oauthInfo.GetID()
|
||||||
newUser := User{
|
newUser := User{
|
||||||
ID: oauthInfo.GetID(),
|
ID: userID,
|
||||||
Username: oauthInfo.Username,
|
Username: oauthInfo.Username,
|
||||||
Nickname: oauthInfo.Name,
|
Nickname: oauthInfo.Name,
|
||||||
Email: oauthInfo.Email,
|
Email: oauthInfo.Email,
|
||||||
@@ -147,6 +161,9 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser
|
|||||||
LastLoginAt: now,
|
LastLoginAt: now,
|
||||||
IsAdmin: false,
|
IsAdmin: false,
|
||||||
}
|
}
|
||||||
|
if err := newUser.assignIDIfMissing(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
if err := tx.Create(&newUser).Error; err != nil {
|
if err := tx.Create(&newUser).Error; err != nil {
|
||||||
return err
|
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 {
|
if err := tx.Create(u).Error; err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -115,7 +115,10 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 生成唯一的 TaskID
|
// 生成唯一的 TaskID
|
||||||
taskID := generateTaskID(taskType, triggeredBy)
|
taskID, err := generateTaskID(taskType, triggeredBy)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
// 创建任务执行记录
|
// 创建任务执行记录
|
||||||
execution := &model.TaskExecution{
|
execution := &model.TaskExecution{
|
||||||
@@ -453,9 +456,12 @@ func handleSuccessfulTask(ctx context.Context, execution *model.TaskExecution, t
|
|||||||
}
|
}
|
||||||
|
|
||||||
// generateTaskID 生成任务 ID
|
// generateTaskID 生成任务 ID
|
||||||
func generateTaskID(taskType string, triggeredBy string) string {
|
func generateTaskID(taskType string, triggeredBy string) (string, error) {
|
||||||
uniqueID := idgen.NextUint64ID()
|
uniqueID, err := idgen.NextUint64ID()
|
||||||
return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, uniqueID)
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s_%s_%d", triggeredBy, taskType, uniqueID), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateRetryTaskID 生成重试任务 ID
|
// generateRetryTaskID 生成重试任务 ID
|
||||||
|
|||||||
@@ -387,8 +387,10 @@ func TestRetryTaskNonExistent(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGenerateTaskID(t *testing.T) {
|
func TestGenerateTaskID(t *testing.T) {
|
||||||
id1 := generateTaskID("test_type", "manual")
|
id1, err := generateTaskID("test_type", "manual")
|
||||||
id2 := generateTaskID("test_type", "manual")
|
require.NoError(t, err)
|
||||||
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user