perf: access token cache

This commit is contained in:
ryan
2026-06-20 09:18:52 +08:00
parent d0a9958711
commit 080be1e03a
17 changed files with 498 additions and 156 deletions
+9 -34
View File
@@ -11,7 +11,6 @@ import (
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/gin-gonic/gin"
@@ -43,8 +42,8 @@ func ListAccessTokens(c *gin.Context) {
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
ctx := c.Request.Context()
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil {
tokens, err := listAccessTokensLogic(ctx, currUser.ID)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -91,8 +90,8 @@ func CreateAccessToken(c *gin.Context) {
maxLimit = val
}
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil {
count, err := countAccessTokensLogic(ctx, currUser.ID)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -120,7 +119,7 @@ func CreateAccessToken(c *gin.Context) {
IsAdmin: req.IsAdmin,
}
if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil {
if err := createAccessTokenLogic(ctx, &tokenRecord); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -152,14 +151,8 @@ func DeleteAccessToken(c *gin.Context) {
return
}
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{})
if tx.Error != nil {
response.AbortBadRequest(c, tx.Error.Error())
return
}
if tx.RowsAffected == 0 {
response.AbortBadRequest(c, errTokenNotFoundOrForbidden)
if err := deleteAccessTokenLogic(ctx, id, currUser.ID); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -187,32 +180,14 @@ func RotateAccessToken(c *gin.Context) {
return
}
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil {
response.AbortBadRequest(c, errTokenNotFoundOrForbidden)
return
}
// 生成新的 Token
newTokenStr, err := model.GenerateTokenString()
newTokenStr, tokenRecord, err := rotateAccessTokenLogic(ctx, id, currUser.ID)
if err != nil {
response.AbortBadRequest(c, errGenerateTokenFailed)
return
}
newTokenHash := model.HashToken(newTokenStr)
newMaskedToken := model.MaskTokenString(newTokenStr)
tokenRecord.TokenHash = newTokenHash
tokenRecord.MaskedToken = newMaskedToken
if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil {
response.AbortBadRequest(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(tokenResponse{
Token: newTokenStr,
Record: tokenRecord,
Record: *tokenRecord,
}))
}
+122
View File
@@ -12,6 +12,7 @@ import (
"math/big"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
@@ -285,3 +286,124 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn
}
return &dbUser, nil
}
func getUserByUsernameOrEmail(ctx context.Context, input string) (*model.User, error) {
var user model.User
if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
func updateLastLogin(ctx context.Context, user *model.User) error {
return db.DB(ctx).Model(user).Update("last_login_at", user.LastLoginAt).Error
}
func registerUserLogic(ctx context.Context, u *model.User) error {
if err := u.RegisterUser(ctx, db.DB(ctx)); err != nil {
if strings.Contains(err.Error(), "duplicate key") || strings.Contains(err.Error(), "UNIQUE") {
return errors.New("用户名或邮箱已被占用")
}
return errors.New("注册失败,请稍后再试")
}
return nil
}
func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass string) error {
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil {
return errors.New(errUserNotFound)
}
if !dbUser.CheckPassword(oldPass) {
return errors.New(errOldPasswordIncorrect)
}
if err := dbUser.SetEncryptedPassword(newPass); err != nil {
return errors.New(errPasswordEncryptFailed)
}
if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil {
return errors.New("更新密码失败,请稍后再试")
}
// 吊销该用户所有的 Access Token
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Find(&tokens).Error; err == nil {
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil {
return errors.New("吊销 Access Token 失败,请稍后再试")
}
oauth.InvalidateCachedUser(ctx, dbUser.ID)
return nil
}
func listAccessTokensLogic(ctx context.Context, userID uint64) ([]model.AccessToken, error) {
var tokens []model.AccessToken
if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil {
return nil, errors.New("获取令牌列表失败,请稍后再试")
}
return tokens, nil
}
func countAccessTokensLogic(ctx context.Context, userID uint64) (int64, error) {
var count int64
if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil {
return 0, errors.New("查询令牌数量失败,请稍后再试")
}
return count, nil
}
func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) error {
if err := db.DB(ctx).Create(record).Error; err != nil {
return errors.New("创建令牌失败,请稍后再试")
}
return nil
}
func deleteAccessTokenLogic(ctx context.Context, id, userID uint64) error {
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil {
return errors.New(errTokenNotFoundOrForbidden)
}
oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash)
tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{})
if tx.Error != nil {
return errors.New("删除令牌失败,请稍后再试")
}
if tx.RowsAffected == 0 {
return errors.New(errTokenNotFoundOrForbidden)
}
return nil
}
func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *model.AccessToken, error) {
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil {
return "", nil, errors.New(errTokenNotFoundOrForbidden)
}
oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash)
newTokenStr, err := model.GenerateTokenString()
if err != nil {
return "", nil, errors.New(errGenerateTokenFailed)
}
newTokenHash := model.HashToken(newTokenStr)
newMaskedToken := model.MaskTokenString(newTokenStr)
tokenRecord.TokenHash = newTokenHash
tokenRecord.MaskedToken = newMaskedToken
if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil {
return "", nil, errors.New("轮换令牌失败,请稍后再试")
}
return newTokenStr, &tokenRecord, nil
}
+18 -35
View File
@@ -13,7 +13,6 @@ import (
"github.com/Rain-kl/Wavelet/internal/common"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -117,8 +116,8 @@ func Login(c *gin.Context) {
return
}
var user model.User
if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil {
user, err := getUserByUsernameOrEmail(ctx, req.Username)
if err != nil {
logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP())
response.AbortBadRequest(c, errUsernameOrPasswordWrong)
return
@@ -139,7 +138,7 @@ func Login(c *gin.Context) {
}
if isEmailLoginVerificationEnabled(ctx) {
result, err := processLoginEmailVerification(ctx, req.Code, &user)
result, err := processLoginEmailVerification(ctx, req.Code, user)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -160,20 +159,20 @@ func Login(c *gin.Context) {
}
user.LastLoginAt = time.Now()
if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil {
response.AbortBadRequest(c, err.Error())
if err := updateLastLogin(ctx, user); err != nil {
response.AbortBadRequest(c, "更新登录时间失败,请稍后再试")
return
}
if err := setLoginSession(ctx, c, &user); err != nil {
if err := setLoginSession(ctx, c, user); err != nil {
response.AbortBadRequest(c, errSaveSessionFailed)
return
}
logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
listener.EmitAdminLoggedIn(ctx, user, c.ClientIP())
c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&user, needChangePassword)))
c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(user, needChangePassword)))
}
// Register 用户注册
@@ -247,7 +246,7 @@ func Register(c *gin.Context) {
return
}
if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil {
if err := registerUserLogic(ctx, &user); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
@@ -275,6 +274,13 @@ func Logout(c *gin.Context) {
username := session.Get(oauth.UserNameKey)
if userID != nil {
logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP())
if id, ok := userID.(uint64); ok {
oauth.InvalidateCachedUser(c.Request.Context(), id)
} else if idFloat, ok := userID.(float64); ok {
oauth.InvalidateCachedUser(c.Request.Context(), uint64(idFloat))
} else if idInt, ok := userID.(int); ok && idInt >= 0 {
oauth.InvalidateCachedUser(c.Request.Context(), uint64(idInt))
}
}
session.Options(oauth.GetSessionOptions(-1))
session.Clear()
@@ -327,35 +333,11 @@ func ChangePassword(c *gin.Context) {
}
ctx := c.Request.Context()
var dbUser model.User
if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil {
response.AbortBadRequest(c, errUserNotFound)
return
}
// 校验旧密码
if !dbUser.CheckPassword(req.OldPassword) {
response.AbortBadRequest(c, errOldPasswordIncorrect)
return
}
// 加密并更新为新密码
if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil {
response.AbortBadRequest(c, errPasswordEncryptFailed)
return
}
if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil {
if err := changePasswordLogic(ctx, userObj.ID, req.OldPassword, req.NewPassword); err != nil {
response.AbortBadRequest(c, err.Error())
return
}
// 吊销该用户所有的 Access Token
if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil {
response.AbortBadRequest(c, "吊销 Access Token 失败: "+err.Error())
return
}
// 销毁当前活跃会话以强制重新登录
session := sessions.Default(c)
session.Clear()
@@ -431,6 +413,7 @@ func UpdateProfile(c *gin.Context) {
response.AbortBadRequest(c, err.Error())
return
}
oauth.InvalidateCachedUser(ctx, userObj.ID)
session := sessions.Default(c)
needChange := session.Get("need_change_password") == true