fix(oauth): avoid overwriting user state when binding

This commit is contained in:
CaIon
2026-08-11 22:03:45 +08:00
parent 3d5dc36f1d
commit d7992672a6
11 changed files with 125 additions and 22 deletions
+1
View File
@@ -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()
+5 -10
View File
@@ -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
+4 -12
View File
@@ -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
}
+24
View File
@@ -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{}{}
+60
View File
@@ -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)
+5
View File
@@ -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"
}
+5
View File
@@ -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
+5
View File
@@ -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"
}
+5
View File
@@ -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
+5
View File
@@ -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"
}
+6
View File
@@ -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
}