mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-02 06:56:36 +08:00
perf: access token cache
This commit is contained in:
@@ -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,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user