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
This commit is contained in:
ryan
2026-06-13 10:27:28 +08:00
parent eb999eba09
commit 26e12594a2
5 changed files with 167 additions and 21 deletions
+1
View File
@@ -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 格式与过期时间
+25 -19
View File
@@ -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
}
+1
View File
@@ -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
+9 -2
View File
@@ -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("密码修改成功"))
+131
View File
@@ -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)
}
}