From d3d7c783a9b07e8f287dfbbe8a22512a407f5158 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 29 Aug 2026 11:49:14 +0800 Subject: [PATCH] fix(auth): synchronize need_change_password across login, user-info and repositories --- backend/plugins/domain/auth/handlers.go | 2 +- backend/plugins/domain/auth/models.go | 2 +- backend/plugins/domain/auth/repository.go | 31 ++++++++++++++++++++--- backend/plugins/domain/user/handlers.go | 3 +++ 4 files changed, 32 insertions(+), 6 deletions(-) diff --git a/backend/plugins/domain/auth/handlers.go b/backend/plugins/domain/auth/handlers.go index 77a9ab9a..c7a8f2a4 100644 --- a/backend/plugins/domain/auth/handlers.go +++ b/backend/plugins/domain/auth/handlers.go @@ -440,7 +440,7 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSou func UserInfo(c *gin.Context) { user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) session := sessions.Default(c) - needChange := session.Get("need_change_password") == true + needChange := session.Get("need_change_password") == true || (user != nil && user.NeedChangePassword) c.JSON( http.StatusOK, diff --git a/backend/plugins/domain/auth/models.go b/backend/plugins/domain/auth/models.go index 87e441c1..74289f2f 100644 --- a/backend/plugins/domain/auth/models.go +++ b/backend/plugins/domain/auth/models.go @@ -172,7 +172,7 @@ func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo Email: user.Email, AvatarURL: user.AvatarURL, IsAdmin: user.IsAdmin, - NeedChangePassword: needChange, + NeedChangePassword: needChange || user.NeedChangePassword, Bio: user.Bio, Phone: user.Phone, Gender: user.Gender, diff --git a/backend/plugins/domain/auth/repository.go b/backend/plugins/domain/auth/repository.go index c245e924..883dc78d 100644 --- a/backend/plugins/domain/auth/repository.go +++ b/backend/plugins/domain/auth/repository.go @@ -8,6 +8,7 @@ import ( "Wavelet/core/contracts" "Wavelet/pkg/util" "context" + "strings" "sync" "time" @@ -79,19 +80,41 @@ func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*CachedToken, // GetActiveUserByID 读取仍处于启用状态的用户 func GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { - var user contracts.UserDTO - if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil { + var row struct { + contracts.UserDTO + Password string `gorm:"column:password"` + } + if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&row).Error; err != nil { return nil, err } + user := row.UserDTO + if row.Password != "" && + !strings.HasPrefix(row.Password, "$2a$") && + !strings.HasPrefix(row.Password, "$2b$") && + !strings.HasPrefix(row.Password, "$2y$") && + !strings.HasPrefix(row.Password, "$2x$") { + user.NeedChangePassword = true + } return &user, nil } // GetUserByID 按 ID 读取用户(不限制启用状态) func GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { - var user contracts.UserDTO - if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil { + var row struct { + contracts.UserDTO + Password string `gorm:"column:password"` + } + if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&row).Error; err != nil { return nil, err } + user := row.UserDTO + if row.Password != "" && + !strings.HasPrefix(row.Password, "$2a$") && + !strings.HasPrefix(row.Password, "$2b$") && + !strings.HasPrefix(row.Password, "$2y$") && + !strings.HasPrefix(row.Password, "$2x$") { + user.NeedChangePassword = true + } return &user, nil } diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index e306cb53..7755f740 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -79,6 +79,9 @@ func Login(c *gin.Context) { sess := sessions.Default(c) sess.Set(contracts.AuthUserIDKey, user.ID) sess.Set(contracts.AuthUserNameKey, user.Username) + needChange := user.NeedChangePassword || user.IsPlaintextPassword() + user.NeedChangePassword = needChange + sess.Set("need_change_password", needChange) if err := sess.Save(); err != nil { logger.ErrorF(c.Request.Context(), "save session failed on login: %v", err) }