diff --git a/controller/auth_flow_test.go b/controller/auth_flow_test.go index 3fe08b42..5917d053 100644 --- a/controller/auth_flow_test.go +++ b/controller/auth_flow_test.go @@ -46,6 +46,7 @@ func (*authFlowTestOAuthProvider) IsUserIDTaken(string) bool func (*authFlowTestOAuthProvider) FillUserByProviderID(*model.User, string) error { return nil } func (*authFlowTestOAuthProvider) SetProviderUserID(*model.User, string) {} func (*authFlowTestOAuthProvider) GetProviderPrefix() string { return "flow_" } +func (*authFlowTestOAuthProvider) ProviderUserIDColumn() string { return "" } func setupAuthFlowControllerTest(t *testing.T) *authFlowTestOAuthProvider { t.Helper() diff --git a/controller/oauth.go b/controller/oauth.go index a477f5b1..4d9725c1 100644 --- a/controller/oauth.go +++ b/controller/oauth.go @@ -263,25 +263,20 @@ func handleOAuthBind(c *gin.Context, provider oauth.Provider, pendingFlow *model return } - user := model.User{Id: pendingFlow.UserId} - err = user.FillUserById() - if err != nil { - common.ApiError(c, err) - return - } + userId := pendingFlow.UserId // Handle binding based on provider type if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok { // Custom provider: use user_oauth_bindings table - err = model.UpdateUserOAuthBinding(user.Id, genericProvider.GetProviderId(), oauthUser.ProviderUserID) + err = model.UpdateUserOAuthBinding(userId, genericProvider.GetProviderId(), oauthUser.ProviderUserID) if err != nil { common.ApiError(c, err) return } } else { - // Built-in provider: update user record directly - provider.SetProviderUserID(&user, oauthUser.ProviderUserID) - err = user.Update(false) + // Built-in provider: 只更新绑定列。完整快照的 user.Update 会把读取时刻的 + // role/status/group 一并写回,覆盖并发发生的封禁、降权或分组变更。 + err = model.UpdateUserBindColumn(userId, provider.ProviderUserIDColumn(), oauthUser.ProviderUserID) if err != nil { common.ApiError(c, err) return diff --git a/controller/wechat.go b/controller/wechat.go index dd185735..f44a2c59 100644 --- a/controller/wechat.go +++ b/controller/wechat.go @@ -156,21 +156,13 @@ func WeChatBind(c *gin.Context) { }) return } - user := model.User{ - Id: c.GetInt("id"), - } - if user.Id == 0 { + userId := c.GetInt("id") + if userId == 0 { c.JSON(http.StatusUnauthorized, gin.H{"success": false, "message": "未登录"}) return } - err = user.FillUserById() - if err != nil { - common.ApiError(c, err) - return - } - user.WeChatId = wechatId - err = user.Update(false) - if err != nil { + // 只更新绑定列,避免完整用户快照覆盖并发的封禁、降权或分组变更。 + if err := model.UpdateUserBindColumn(userId, "wechat_id", wechatId); err != nil { common.ApiError(c, err) return } diff --git a/model/user.go b/model/user.go index eb4ea086..83d7aeec 100644 --- a/model/user.go +++ b/model/user.go @@ -190,6 +190,30 @@ func UpdateUserSetting(userId int, setting dto.UserSetting) error { return updateUserSettingCache(userId, settingValue) } +// userBindColumns 允许通过 UpdateUserBindColumn 更新的第三方账号绑定列白名单。 +// 列名只可能来自代码内部的 provider 实现,白名单是防御纵深,不依赖调用方自律。 +var userBindColumns = map[string]bool{ + "github_id": true, + "discord_id": true, + "oidc_id": true, + "linux_do_id": true, + "wechat_id": true, +} + +// UpdateUserBindColumn 第三方账号绑定字段的专用更新。 +// 绑定操作必须只写绑定列:若改为“读取完整用户 → 改一个字段 → 整体更新”, +// 读快照期间并发发生的封禁、降权或分组变更会被旧快照覆盖恢复。 +// 角色、状态、分组只允许通过各自带锁/CAS 的专用方法修改。 +func UpdateUserBindColumn(userId int, column string, value string) error { + if userId <= 0 { + return errors.New("id 为空!") + } + if !userBindColumns[column] { + return fmt.Errorf("invalid user bind column: %s", column) + } + return DB.Model(&User{}).Where("id = ?", userId).Update(column, value).Error +} + // 根据用户角色生成默认的边栏配置 func generateDefaultSidebarConfigForRole(userRole int) string { defaultConfig := map[string]interface{}{} diff --git a/model/user_update_test.go b/model/user_update_test.go index 8e69ba1e..1a89e93a 100644 --- a/model/user_update_test.go +++ b/model/user_update_test.go @@ -27,6 +27,21 @@ func setupUserUpdateTestState(t *testing.T) { }) } +func createUserBindTestUser(t *testing.T) User { + t.Helper() + user := User{ + Username: "bind-test-user", + Password: "unused-password-hash", + Role: common.RoleCommonUser, + Status: common.UserStatusEnabled, + Group: "default", + AuthVersion: 1, + AffCode: "bind-test-aff-code", + } + require.NoError(t, DB.Create(&user).Error) + return user +} + func TestUserUpdateDoesNotOverwriteConcurrentAccountingOrTokenChanges(t *testing.T) { setupUserUpdateTestState(t) @@ -218,6 +233,51 @@ func TestInsertKeepsBlankPasswordForPasswordlessUser(t *testing.T) { assert.Empty(t, stored.Password) } +func TestUpdateUserBindColumnOnlyTouchesTheBindingColumn(t *testing.T) { + truncateTables(t) + + user := createUserBindTestUser(t) + require.NoError(t, DB.Model(&User{}).Where("id = ?", user.Id).Updates(map[string]interface{}{ + "role": common.RoleAdminUser, + "status": common.UserStatusEnabled, + "group": "vip", + }).Error) + + require.NoError(t, UpdateUserBindColumn(user.Id, "github_id", "gh-12345")) + + reloaded, err := GetUserById(user.Id, true) + require.NoError(t, err) + assert.Equal(t, "gh-12345", reloaded.GitHubId) + assert.Equal(t, common.RoleAdminUser, reloaded.Role) + assert.Equal(t, common.UserStatusEnabled, reloaded.Status) + assert.Equal(t, "vip", reloaded.Group) +} + +func TestUpdateUserBindColumnPreservesRestrictiveChange(t *testing.T) { + truncateTables(t) + + user := createUserBindTestUser(t) + require.NoError(t, DB.Model(&User{}).Where("id = ?", user.Id). + Update("status", common.UserStatusDisabled).Error) + require.NoError(t, UpdateUserBindColumn(user.Id, "wechat_id", "wx-open-id")) + + reloaded, err := GetUserById(user.Id, true) + require.NoError(t, err) + assert.Equal(t, "wx-open-id", reloaded.WeChatId) + assert.Equal(t, common.UserStatusDisabled, reloaded.Status) +} + +func TestUpdateUserBindColumnRejectsNonWhitelistedColumns(t *testing.T) { + truncateTables(t) + + user := createUserBindTestUser(t) + for _, column := range []string{"role", "status", "group", "quota", "username", "password", "id"} { + assert.Error(t, UpdateUserBindColumn(user.Id, column, "1"), "column %s must be rejected", column) + } + assert.Error(t, UpdateUserBindColumn(user.Id, "github_id; DROP TABLE users", "x")) + assert.Error(t, UpdateUserBindColumn(0, "github_id", "x")) +} + func TestValidateAndFillRejectsPasswordlessUser(t *testing.T) { setupUserUpdateTestState(t) diff --git a/oauth/discord.go b/oauth/discord.go index b626d2f8..7f1ce9a6 100644 --- a/oauth/discord.go +++ b/oauth/discord.go @@ -170,3 +170,8 @@ func (p *DiscordProvider) SetProviderUserID(user *model.User, providerUserID str func (p *DiscordProvider) GetProviderPrefix() string { return "discord_" } + +// ProviderUserIDColumn returns the users-table column storing this provider's user ID. +func (p *DiscordProvider) ProviderUserIDColumn() string { + return "discord_id" +} diff --git a/oauth/generic.go b/oauth/generic.go index 11bbb9b6..23ae2273 100644 --- a/oauth/generic.go +++ b/oauth/generic.go @@ -312,6 +312,11 @@ func (p *GenericOAuthProvider) GetProviderPrefix() string { return p.config.Slug + "_" } +// ProviderUserIDColumn returns the users-table column storing this provider's user ID. +func (p *GenericOAuthProvider) ProviderUserIDColumn() string { + return "" +} + // GetProviderId returns the provider ID for binding purposes func (p *GenericOAuthProvider) GetProviderId() int { return p.config.Id diff --git a/oauth/github.go b/oauth/github.go index 314118a3..524e0fdd 100644 --- a/oauth/github.go +++ b/oauth/github.go @@ -176,3 +176,8 @@ func (p *GitHubProvider) SetProviderUserID(user *model.User, providerUserID stri func (p *GitHubProvider) GetProviderPrefix() string { return "github_" } + +// ProviderUserIDColumn returns the users-table column storing this provider's user ID. +func (p *GitHubProvider) ProviderUserIDColumn() string { + return "github_id" +} diff --git a/oauth/linuxdo.go b/oauth/linuxdo.go index 1ed91e00..cba54b83 100644 --- a/oauth/linuxdo.go +++ b/oauth/linuxdo.go @@ -184,6 +184,11 @@ func (p *LinuxDOProvider) GetProviderPrefix() string { return "linuxdo_" } +// ProviderUserIDColumn returns the users-table column storing this provider's user ID. +func (p *LinuxDOProvider) ProviderUserIDColumn() string { + return "linux_do_id" +} + // TrustLevelError indicates the user's trust level is too low type TrustLevelError struct { Required int diff --git a/oauth/oidc.go b/oauth/oidc.go index 25874329..79582a4e 100644 --- a/oauth/oidc.go +++ b/oauth/oidc.go @@ -175,3 +175,8 @@ func (p *OIDCProvider) SetProviderUserID(user *model.User, providerUserID string func (p *OIDCProvider) GetProviderPrefix() string { return "oidc_" } + +// ProviderUserIDColumn returns the users-table column storing this provider's user ID. +func (p *OIDCProvider) ProviderUserIDColumn() string { + return "oidc_id" +} diff --git a/oauth/provider.go b/oauth/provider.go index 785ed25d..902f50d6 100644 --- a/oauth/provider.go +++ b/oauth/provider.go @@ -33,4 +33,10 @@ type Provider interface { // GetProviderPrefix returns the prefix for auto-generated usernames (e.g., "github_") GetProviderPrefix() string + + // ProviderUserIDColumn returns the users-table column that stores this provider's + // user ID, used by bind flows to update only the binding column instead of + // writing back a full user snapshot. Providers that persist bindings elsewhere + // (e.g. the user_oauth_bindings table) return an empty string. + ProviderUserIDColumn() string }