From dfc0d6324b40c1d6c2972e524409f933541bfb0f Mon Sep 17 00:00:00 2001 From: Calcium-Ion Date: Fri, 3 Jul 2026 15:25:33 +0800 Subject: [PATCH] Merge commit from fork * Harden user setting cache updates * Fix user update test isolation --- controller/subscription.go | 3 +- controller/user.go | 19 +++----- model/user.go | 40 ++++++++++++----- model/user_cache.go | 34 ++++++++++++-- model/user_update_test.go | 92 ++++++++++++++++++++++++++++++++++++++ router/api-router.go | 2 +- 6 files changed, 161 insertions(+), 29 deletions(-) create mode 100644 model/user_update_test.go diff --git a/controller/subscription.go b/controller/subscription.go index 8f53e5b8..7007ce45 100644 --- a/controller/subscription.go +++ b/controller/subscription.go @@ -89,8 +89,7 @@ func UpdateSubscriptionPreference(c *gin.Context) { } current := user.GetSetting() current.BillingPreference = pref - user.SetSetting(current) - if err := user.Update(false); err != nil { + if err := model.UpdateUserSetting(user.Id, current); err != nil { common.ApiError(c, err) return } diff --git a/controller/user.go b/controller/user.go index 1fc52dd9..0e3dc89c 100644 --- a/controller/user.go +++ b/controller/user.go @@ -732,8 +732,7 @@ func AdminClearUserBinding(c *gin.Context) { func UpdateSelf(c *gin.Context) { var requestData map[string]interface{} - err := json.NewDecoder(c.Request.Body).Decode(&requestData) - if err != nil { + if err := common.DecodeJson(c.Request.Body, &requestData); err != nil { common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } @@ -755,9 +754,7 @@ func UpdateSelf(c *gin.Context) { currentSetting.SidebarModules = sidebarModulesStr } - // 保存更新后的设置 - user.SetSetting(currentSetting) - if err := user.Update(false); err != nil { + if err := model.UpdateUserSetting(user.Id, currentSetting); err != nil { common.ApiErrorI18n(c, i18n.MsgUpdateFailed) return } @@ -783,9 +780,7 @@ func UpdateSelf(c *gin.Context) { currentSetting.Language = langStr } - // 保存更新后的设置 - user.SetSetting(currentSetting) - if err := user.Update(false); err != nil { + if err := model.UpdateUserSetting(user.Id, currentSetting); err != nil { common.ApiErrorI18n(c, i18n.MsgUpdateFailed) return } @@ -796,13 +791,12 @@ func UpdateSelf(c *gin.Context) { // 原有的用户信息更新逻辑 var user model.User - requestDataBytes, err := json.Marshal(requestData) + requestDataBytes, err := common.Marshal(requestData) if err != nil { common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } - err = json.Unmarshal(requestDataBytes, &user) - if err != nil { + if err = common.Unmarshal(requestDataBytes, &user); err != nil { common.ApiErrorI18n(c, i18n.MsgInvalidParams) return } @@ -1434,8 +1428,7 @@ func UpdateUserSetting(c *gin.Context) { } // 更新用户设置 - user.SetSetting(settings) - if err := user.Update(false); err != nil { + if err := model.UpdateUserSetting(user.Id, settings); err != nil { common.ApiErrorI18n(c, i18n.MsgUpdateFailed) return } diff --git a/model/user.go b/model/user.go index cef9a726..2f438e87 100644 --- a/model/user.go +++ b/model/user.go @@ -2,7 +2,6 @@ package model import ( "database/sql" - "encoding/json" "errors" "fmt" "strconv" @@ -83,7 +82,7 @@ func (user *User) SetAccessToken(token string) { func (user *User) GetSetting() dto.UserSetting { setting := dto.UserSetting{} if user.Setting != "" { - err := json.Unmarshal([]byte(user.Setting), &setting) + err := common.Unmarshal([]byte(user.Setting), &setting) if err != nil { common.SysLog("failed to unmarshal setting: " + err.Error()) } @@ -92,7 +91,7 @@ func (user *User) GetSetting() dto.UserSetting { } func (user *User) SetSetting(setting dto.UserSetting) { - settingBytes, err := json.Marshal(setting) + settingBytes, err := common.Marshal(setting) if err != nil { common.SysLog("failed to marshal setting: " + err.Error()) return @@ -100,6 +99,21 @@ func (user *User) SetSetting(setting dto.UserSetting) { user.Setting = string(settingBytes) } +func UpdateUserSetting(userId int, setting dto.UserSetting) error { + if userId == 0 { + return errors.New("id 为空!") + } + settingBytes, err := common.Marshal(setting) + if err != nil { + return err + } + settingValue := string(settingBytes) + if err = DB.Model(&User{}).Where("id = ?", userId).Update("setting", settingValue).Error; err != nil { + return err + } + return updateUserSettingCache(userId, settingValue) +} + // 根据用户角色生成默认的边栏配置 func generateDefaultSidebarConfigForRole(userRole int) string { defaultConfig := map[string]interface{}{} @@ -153,7 +167,7 @@ func generateDefaultSidebarConfigForRole(userRole int) string { // 普通用户不包含admin区域 // 转换为JSON字符串 - configBytes, err := json.Marshal(defaultConfig) + configBytes, err := common.Marshal(defaultConfig) if err != nil { common.SysLog("生成默认边栏配置失败: " + err.Error()) return "" @@ -524,11 +538,14 @@ func (user *User) UpdateWithTx(tx *gorm.DB, updatePassword bool) error { } } newUser := *user - tx.First(&user, user.Id) - if err = tx.Model(user).Updates(newUser).Error; err != nil { + current := User{} + if err = tx.First(¤t, user.Id).Error; err != nil { return err } - return nil + if err = tx.Model(¤t).Omit("quota", "used_quota", "request_count").Updates(newUser).Error; err != nil { + return err + } + return tx.First(user, user.Id).Error } func (user *User) Edit(updatePassword bool) error { @@ -558,11 +575,14 @@ func (user *User) EditWithTx(tx *gorm.DB, updatePassword bool) error { updates["password"] = newUser.Password } - tx.First(&user, user.Id) - if err = tx.Model(user).Updates(updates).Error; err != nil { + current := User{} + if err = tx.First(¤t, user.Id).Error; err != nil { return err } - return nil + if err = tx.Model(¤t).Updates(updates).Error; err != nil { + return err + } + return tx.First(user, user.Id).Error } func (user *User) ClearBinding(bindingType string) error { diff --git a/model/user_cache.go b/model/user_cache.go index 80d0264f..2a246c84 100644 --- a/model/user_cache.go +++ b/model/user_cache.go @@ -63,8 +63,7 @@ func InvalidateUserCache(userId int) error { return invalidateUserCache(userId) } -// updateUserCache updates all user cache fields using hash -func updateUserCache(user User) error { +func populateUserCache(user User) error { if !common.RedisEnabled { return nil } @@ -76,6 +75,28 @@ func updateUserCache(user User) error { ) } +// updateUserCache refreshes non-quota user cache fields. +// Quota is maintained by atomic quota delta paths and must not be overwritten +// by stale user snapshots from profile/settings updates. +func updateUserCache(user User) error { + if !common.RedisEnabled { + return nil + } + if err := updateUserGroupCache(user.Id, user.Group); err != nil { + return err + } + if err := updateUserEmailCache(user.Id, user.Email); err != nil { + return err + } + if err := updateUserStatusCache(user.Id, user.Status == common.UserStatusEnabled); err != nil { + return err + } + if err := updateUserNameCache(user.Id, user.Username); err != nil { + return err + } + return updateUserSettingCache(user.Id, user.Setting) +} + // GetUserCache gets complete user cache from hash func GetUserCache(userId int) (userCache *UserBase, err error) { var user *User @@ -84,7 +105,7 @@ func GetUserCache(userId int) (userCache *UserBase, err error) { // Update Redis cache asynchronously on successful DB read if shouldUpdateRedis(fromDB, err) && user != nil { gopool.Go(func() { - if err := updateUserCache(*user); err != nil { + if err := populateUserCache(*user); err != nil { common.SysLog("failed to update user status cache: " + err.Error()) } }) @@ -214,6 +235,13 @@ func UpdateUserGroupCache(userId int, group string) error { return updateUserGroupCache(userId, group) } +func updateUserEmailCache(userId int, email string) error { + if !common.RedisEnabled { + return nil + } + return common.RedisHSetField(getUserCacheKey(userId), "Email", email) +} + func updateUserNameCache(userId int, username string) error { if !common.RedisEnabled { return nil diff --git a/model/user_update_test.go b/model/user_update_test.go new file mode 100644 index 00000000..b04f0c6e --- /dev/null +++ b/model/user_update_test.go @@ -0,0 +1,92 @@ +package model + +import ( + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/dto" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupUserUpdateTestState(t *testing.T) { + t.Helper() + truncateTables(t) + require.NoError(t, DB.Exec("DELETE FROM users").Error) + + oldRedisEnabled := common.RedisEnabled + oldBatchUpdateEnabled := common.BatchUpdateEnabled + common.RedisEnabled = false + common.BatchUpdateEnabled = false + t.Cleanup(func() { + common.RedisEnabled = oldRedisEnabled + common.BatchUpdateEnabled = oldBatchUpdateEnabled + }) +} + +func TestUserUpdateDoesNotOverwriteAccountingFields(t *testing.T) { + setupUserUpdateTestState(t) + + user := User{ + Id: 1, + Username: "quota-race-user", + Password: "password", + DisplayName: "before", + Status: common.UserStatusEnabled, + Quota: 1000, + UsedQuota: 20, + RequestCount: 3, + } + require.NoError(t, DB.Create(&user).Error) + + staleUser, err := GetUserById(user.Id, true) + require.NoError(t, err) + + require.NoError(t, DB.Model(&User{}).Where("id = ?", user.Id).Updates(map[string]interface{}{ + "quota": gorm.Expr("quota - ?", 400), + "used_quota": gorm.Expr("used_quota + ?", 400), + "request_count": gorm.Expr("request_count + ?", 1), + }).Error) + + staleUser.DisplayName = "after" + require.NoError(t, staleUser.Update(false)) + + var got User + require.NoError(t, DB.First(&got, user.Id).Error) + assert.Equal(t, "after", got.DisplayName) + assert.Equal(t, 600, got.Quota) + assert.Equal(t, 420, got.UsedQuota) + assert.Equal(t, 4, got.RequestCount) +} + +func TestUpdateUserSettingOnlyUpdatesSetting(t *testing.T) { + setupUserUpdateTestState(t) + + user := User{ + Id: 2, + Username: "setting-user", + Password: "password", + Status: common.UserStatusEnabled, + Quota: 1000, + UsedQuota: 20, + RequestCount: 3, + } + require.NoError(t, DB.Create(&user).Error) + + require.NoError(t, DB.Model(&User{}).Where("id = ?", user.Id).Updates(map[string]interface{}{ + "quota": gorm.Expr("quota - ?", 250), + "used_quota": gorm.Expr("used_quota + ?", 250), + "request_count": gorm.Expr("request_count + ?", 1), + }).Error) + + require.NoError(t, UpdateUserSetting(user.Id, dto.UserSetting{Language: "zh"})) + + var got User + require.NoError(t, DB.First(&got, user.Id).Error) + assert.Equal(t, 750, got.Quota) + assert.Equal(t, 270, got.UsedQuota) + assert.Equal(t, 4, got.RequestCount) + assert.Equal(t, "zh", got.GetSetting().Language) +} diff --git a/router/api-router.go b/router/api-router.go index efe2131d..86726eed 100644 --- a/router/api-router.go +++ b/router/api-router.go @@ -83,7 +83,7 @@ func SetApiRouter(router *gin.Engine) { selfRoute.GET("/self/groups", controller.GetUserGroups) selfRoute.GET("/self", controller.GetSelf) selfRoute.GET("/models", controller.GetUserModels) - selfRoute.PUT("/self", controller.UpdateSelf) + selfRoute.PUT("/self", middleware.CriticalRateLimit(), controller.UpdateSelf) selfRoute.DELETE("/self", controller.DeleteSelf) selfRoute.GET("/token", controller.GenerateAccessToken) selfRoute.GET("/passkey", controller.PasskeyStatus)