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) FillUserByProviderID(*model.User, string) error { return nil }
|
||||||
func (*authFlowTestOAuthProvider) SetProviderUserID(*model.User, string) {}
|
func (*authFlowTestOAuthProvider) SetProviderUserID(*model.User, string) {}
|
||||||
func (*authFlowTestOAuthProvider) GetProviderPrefix() string { return "flow_" }
|
func (*authFlowTestOAuthProvider) GetProviderPrefix() string { return "flow_" }
|
||||||
|
func (*authFlowTestOAuthProvider) ProviderUserIDColumn() string { return "" }
|
||||||
|
|
||||||
func setupAuthFlowControllerTest(t *testing.T) *authFlowTestOAuthProvider {
|
func setupAuthFlowControllerTest(t *testing.T) *authFlowTestOAuthProvider {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|||||||
+5
-10
@@ -263,25 +263,20 @@ func handleOAuthBind(c *gin.Context, provider oauth.Provider, pendingFlow *model
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
user := model.User{Id: pendingFlow.UserId}
|
userId := pendingFlow.UserId
|
||||||
err = user.FillUserById()
|
|
||||||
if err != nil {
|
|
||||||
common.ApiError(c, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handle binding based on provider type
|
// Handle binding based on provider type
|
||||||
if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok {
|
if genericProvider, ok := provider.(*oauth.GenericOAuthProvider); ok {
|
||||||
// Custom provider: use user_oauth_bindings table
|
// 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 {
|
if err != nil {
|
||||||
common.ApiError(c, err)
|
common.ApiError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
// Built-in provider: update user record directly
|
// Built-in provider: 只更新绑定列。完整快照的 user.Update 会把读取时刻的
|
||||||
provider.SetProviderUserID(&user, oauthUser.ProviderUserID)
|
// role/status/group 一并写回,覆盖并发发生的封禁、降权或分组变更。
|
||||||
err = user.Update(false)
|
err = model.UpdateUserBindColumn(userId, provider.ProviderUserIDColumn(), oauthUser.ProviderUserID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
common.ApiError(c, err)
|
common.ApiError(c, err)
|
||||||
return
|
return
|
||||||
|
|||||||
+4
-12
@@ -156,21 +156,13 @@ func WeChatBind(c *gin.Context) {
|
|||||||
})
|
})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
user := model.User{
|
userId := c.GetInt("id")
|
||||||
Id: c.GetInt("id"),
|
if userId == 0 {
|
||||||
}
|
|
||||||
if user.Id == 0 {
|
|
||||||
c.JSON(http.StatusUnauthorized, gin.H{"success": false, "message": "未登录"})
|
c.JSON(http.StatusUnauthorized, gin.H{"success": false, "message": "未登录"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = user.FillUserById()
|
// 只更新绑定列,避免完整用户快照覆盖并发的封禁、降权或分组变更。
|
||||||
if err != nil {
|
if err := model.UpdateUserBindColumn(userId, "wechat_id", wechatId); err != nil {
|
||||||
common.ApiError(c, err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
user.WeChatId = wechatId
|
|
||||||
err = user.Update(false)
|
|
||||||
if err != nil {
|
|
||||||
common.ApiError(c, err)
|
common.ApiError(c, err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -190,6 +190,30 @@ func UpdateUserSetting(userId int, setting dto.UserSetting) error {
|
|||||||
return updateUserSettingCache(userId, settingValue)
|
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 {
|
func generateDefaultSidebarConfigForRole(userRole int) string {
|
||||||
defaultConfig := map[string]interface{}{}
|
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) {
|
func TestUserUpdateDoesNotOverwriteConcurrentAccountingOrTokenChanges(t *testing.T) {
|
||||||
setupUserUpdateTestState(t)
|
setupUserUpdateTestState(t)
|
||||||
|
|
||||||
@@ -218,6 +233,51 @@ func TestInsertKeepsBlankPasswordForPasswordlessUser(t *testing.T) {
|
|||||||
assert.Empty(t, stored.Password)
|
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) {
|
func TestValidateAndFillRejectsPasswordlessUser(t *testing.T) {
|
||||||
setupUserUpdateTestState(t)
|
setupUserUpdateTestState(t)
|
||||||
|
|
||||||
|
|||||||
@@ -170,3 +170,8 @@ func (p *DiscordProvider) SetProviderUserID(user *model.User, providerUserID str
|
|||||||
func (p *DiscordProvider) GetProviderPrefix() string {
|
func (p *DiscordProvider) GetProviderPrefix() string {
|
||||||
return "discord_"
|
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 + "_"
|
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
|
// GetProviderId returns the provider ID for binding purposes
|
||||||
func (p *GenericOAuthProvider) GetProviderId() int {
|
func (p *GenericOAuthProvider) GetProviderId() int {
|
||||||
return p.config.Id
|
return p.config.Id
|
||||||
|
|||||||
@@ -176,3 +176,8 @@ func (p *GitHubProvider) SetProviderUserID(user *model.User, providerUserID stri
|
|||||||
func (p *GitHubProvider) GetProviderPrefix() string {
|
func (p *GitHubProvider) GetProviderPrefix() string {
|
||||||
return "github_"
|
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_"
|
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
|
// TrustLevelError indicates the user's trust level is too low
|
||||||
type TrustLevelError struct {
|
type TrustLevelError struct {
|
||||||
Required int
|
Required int
|
||||||
|
|||||||
@@ -175,3 +175,8 @@ func (p *OIDCProvider) SetProviderUserID(user *model.User, providerUserID string
|
|||||||
func (p *OIDCProvider) GetProviderPrefix() string {
|
func (p *OIDCProvider) GetProviderPrefix() string {
|
||||||
return "oidc_"
|
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 returns the prefix for auto-generated usernames (e.g., "github_")
|
||||||
GetProviderPrefix() string
|
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