fix(oauth): avoid overwriting user state when binding
This commit is contained in:
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{}{}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user