From 26e12594a23b56ece56f096c7fe978c8d56e5924 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 13 Jun 2026 10:27:28 +0800 Subject: [PATCH] fix(user): revoke all sessions and access tokens on password change - Store user password hash in session during login - Validate password hash compatibility on requests to prevent session reuse - Revoke all user access tokens and clear session on ChangePassword --- internal/apps/oauth/constants.go | 1 + internal/apps/oauth/middlewares.go | 44 +++++----- internal/apps/oauth/sources.go | 1 + internal/apps/user/routers.go | 11 ++- internal/apps/user/routers_test.go | 131 +++++++++++++++++++++++++++++ 5 files changed, 167 insertions(+), 21 deletions(-) diff --git a/internal/apps/oauth/constants.go b/internal/apps/oauth/constants.go index bb066a61..64231d67 100644 --- a/internal/apps/oauth/constants.go +++ b/internal/apps/oauth/constants.go @@ -17,6 +17,7 @@ const ( TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权 TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限 SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials + PasswordHashKey = "password_hash" ) // OAuth State 缓存 Key 格式与过期时间 diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index aecc04cf..84e0c785 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -13,6 +13,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/otel_trace" "github.com/Rain-kl/Wavelet/internal/util" + "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" ) @@ -41,39 +42,44 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { } var user model.User - var authenticated bool - var tokenAuth bool - var tokenAdmin bool + // 优先使用 Access Token 鉴权 if tokenStr != "" { tokenHash := model.HashToken(tokenStr) var tokenRecord model.AccessToken if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err == nil { if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err == nil { - authenticated = true - tokenAuth = true - tokenAdmin = tokenRecord.IsAdmin + util.SetToContext(c, TokenAuthKey, true) + util.SetToContext(c, TokenAdminKey, tokenRecord.IsAdmin) + return &user, nil } } } - if !authenticated { - // load user from session - userID := GetUserIDFromContext(c) - if userID <= 0 { - return nil, errors.New("unauthorized") - } + // 降级使用 Session 鉴权 + userID := GetUserIDFromContext(c) + if userID <= 0 { + return nil, errors.New("unauthorized") + } - // load user from db to make sure is active - tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&user) - if tx.Error != nil { - return nil, tx.Error + // load user from db to make sure is active + tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&user) + if tx.Error != nil { + return nil, tx.Error + } + + // 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致 + if user.Password != "" { + session := sessions.Default(c) + sessionHash, _ := session.Get(PasswordHashKey).(string) + if sessionHash != user.Password { + return nil, errors.New("session expired due to password change") } } - // set keys in context - util.SetToContext(c, TokenAuthKey, tokenAuth) - util.SetToContext(c, TokenAdminKey, tokenAdmin) + // set keys in context for session auth + util.SetToContext(c, TokenAuthKey, false) + util.SetToContext(c, TokenAdminKey, false) return &user, nil } diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index 5651580d..cf7472e7 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -197,6 +197,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro session := sessions.Default(c) session.Set(UserIDKey, user.ID) session.Set(UserNameKey, user.Username) + session.Set(PasswordHashKey, user.Password) // 根据系统配置动态设置 Session 过期时间 maxAge := config.Config.App.SessionAge diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index 1e632a9d..6458b1ff 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -63,6 +63,7 @@ func setLoginSession(ctx context.Context, c *gin.Context, user *model.User) erro session := sessions.Default(c) session.Set(oauth.UserIDKey, user.ID) session.Set(oauth.UserNameKey, user.Username) + session.Set(oauth.PasswordHashKey, user.Password) // 根据系统配置动态设置 Session 过期时间 maxAge := config.Config.App.SessionAge @@ -337,9 +338,15 @@ func ChangePassword(c *gin.Context) { return } - // 清除 Session 中修改密码提示状态 + // 吊销该用户所有的 Access Token + if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { + c.JSON(http.StatusOK, util.Err("吊销 Access Token 失败: "+err.Error())) + return + } + + // 销毁当前活跃会话以强制重新登录 session := sessions.Default(c) - session.Delete("need_change_password") + session.Clear() _ = session.Save() c.JSON(http.StatusOK, util.OK("密码修改成功")) diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index 3bc2055a..7e305422 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -9,6 +9,7 @@ import ( "encoding/json" "net/http" "net/http/httptest" + "strings" "testing" "time" @@ -536,6 +537,7 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) { session := sessions.Default(c) session.Set(oauth.UserIDKey, userID) session.Set(oauth.UserNameKey, username) + session.Set(oauth.PasswordHashKey, userRecord.Password) _ = session.Save() c.String(http.StatusOK, "ok") }) @@ -555,3 +557,132 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) { } } +func TestChangePasswordRevocation(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + // 1. Seed a user with a password + user := model.User{ + ID: uint64(888), + Username: "revoketest", + Nickname: "Revoke Test", + IsActive: true, + } + if err := user.SetEncryptedPassword("oldpassword123"); err != nil { + t.Fatalf("set encrypted password failed: %v", err) + } + if err := dbConn.Create(&user).Error; err != nil { + t.Fatalf("create test user failed: %v", err) + } + + // 2. Seed an active AccessToken for this user + tokenStr, err := model.GenerateTokenString() + if err != nil { + t.Fatalf("generate token string failed: %v", err) + } + tokenHash := model.HashToken(tokenStr) + tokenRecord := model.AccessToken{ + UserID: user.ID, + Name: "Test Token", + TokenHash: tokenHash, + MaskedToken: model.MaskTokenString(tokenStr), + } + if err := dbConn.Create(&tokenRecord).Error; err != nil { + t.Fatalf("create test access token failed: %v", err) + } + + // 3. Set up router + gin.SetMode(gin.TestMode) + r := gin.New() + store := cookie.NewStore([]byte("test_session_secret")) + r.Use(sessions.Sessions("test_session_id", store)) + + r.GET("/mock-login", func(c *gin.Context) { + session := sessions.Default(c) + session.Set(oauth.UserIDKey, user.ID) + session.Set(oauth.UserNameKey, user.Username) + session.Set(oauth.PasswordHashKey, user.Password) + _ = session.Save() + c.String(http.StatusOK, "ok") + }) + + r.GET("/mock-old-session-login", func(c *gin.Context) { + session := sessions.Default(c) + session.Set(oauth.UserIDKey, user.ID) + session.Set(oauth.UserNameKey, user.Username) + session.Set(oauth.PasswordHashKey, "invalid_old_password_hash") + _ = session.Save() + c.String(http.StatusOK, "ok") + }) + + r.POST("/api/v1/user/change-password", oauth.LoginRequired(), ChangePassword) + r.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), ListAccessTokens) + + // 4. Perform mock login to get cookie + wMock := httptest.NewRecorder() + reqMock, _ := http.NewRequest(http.MethodGet, "/mock-login", nil) + r.ServeHTTP(wMock, reqMock) + cookieVal := wMock.Header().Get("Set-Cookie") + + // 5. Test that session and token work initially + reqSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) + reqSession.Header.Set("Cookie", cookieVal) + wSession := httptest.NewRecorder() + r.ServeHTTP(wSession, reqSession) + if wSession.Code != http.StatusOK { + t.Errorf("expected 200, got %d", wSession.Code) + } + + reqToken, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) + reqToken.Header.Set("X-Access-Token", tokenStr) + wToken := httptest.NewRecorder() + r.ServeHTTP(wToken, reqToken) + if wToken.Code != http.StatusOK { + t.Errorf("expected 200 for token, got %d", wToken.Code) + } + + // 6. Change password using the active session + reqBody := `{"old_password": "oldpassword123", "new_password": "newpassword12345"}` + reqChange, _ := http.NewRequest(http.MethodPost, "/api/v1/user/change-password", strings.NewReader(reqBody)) + reqChange.Header.Set("Content-Type", "application/json") + reqChange.Header.Set("Cookie", cookieVal) + wChange := httptest.NewRecorder() + r.ServeHTTP(wChange, reqChange) + if wChange.Code != http.StatusOK { + t.Fatalf("expected change password to return 200, got %d. Body: %s", wChange.Code, wChange.Body.String()) + } + + // 7. Verification: The active session that performed change-password is now cleared (401) + reqSessionAfter, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) + reqSessionAfter.Header.Set("Cookie", cookieVal) + wSessionAfter := httptest.NewRecorder() + r.ServeHTTP(wSessionAfter, reqSessionAfter) + if wSessionAfter.Code != http.StatusUnauthorized { + t.Errorf("expected session to be revoked (401), got %d", wSessionAfter.Code) + } + + // 8. Verification: An old session (holding an outdated password hash) should be rejected (401) + wMockOld := httptest.NewRecorder() + reqMockOld, _ := http.NewRequest(http.MethodGet, "/mock-old-session-login", nil) + r.ServeHTTP(wMockOld, reqMockOld) + oldCookieVal := wMockOld.Header().Get("Set-Cookie") + + reqOldSession, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) + reqOldSession.Header.Set("Cookie", oldCookieVal) + wOldSession := httptest.NewRecorder() + r.ServeHTTP(wOldSession, reqOldSession) + if wOldSession.Code != http.StatusUnauthorized { + t.Errorf("expected old session with invalid hash to return 401, got %d", wOldSession.Code) + } + + // 9. Verification: The Access Token should be deleted from DB and rejected (401) + reqTokenAfter, _ := http.NewRequest(http.MethodGet, "/api/v1/user/access-tokens", nil) + reqTokenAfter.Header.Set("X-Access-Token", tokenStr) + wTokenAfter := httptest.NewRecorder() + r.ServeHTTP(wTokenAfter, reqTokenAfter) + if wTokenAfter.Code != http.StatusUnauthorized { + t.Errorf("expected access token to be revoked (401), got %d", wTokenAfter.Code) + } +} + +