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
@@ -13,6 +13,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
"github.com/Rain-kl/Wavelet/internal/listener"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/hibiken/asynq"
@@ -88,7 +89,7 @@ func enableAdminLoginEvent(t *testing.T, dbConn *gorm.DB, channelName string, ta
event.Enabled = true
event.Channels = []string{channelName}
event.Targets = targets
require.NoError(t, dbConn.Save(&event).Error)
require.NoError(t, repository.SavePushEvent(context.Background(), &event))
}
func waitForAsyncTrigger(t *testing.T) {
@@ -165,11 +166,11 @@ func TestAdminLoginPushIntegration(t *testing.T) {
var event model.PushEvent
require.NoError(t, dbConn.Where("event_key = ?", AdminLogin.Key).First(&event).Error)
event.Enabled = false
require.NoError(t, dbConn.Save(&event).Error)
require.NoError(t, repository.SavePushEvent(context.Background(), &event))
listener.EmitAdminLoggedIn(context.Background(), adminUser, "10.0.0.1")
waitForAsyncTrigger(t)
assert.Equal(t, int64(0), countPushTasks(t, dbConn))
})
}
}
+10 -7
View File
@@ -3,7 +3,8 @@
package push
import ("bytes"
import (
"bytes"
"context"
"encoding/json"
"net/http"
@@ -13,7 +14,10 @@ import ("bytes"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
pkgpush "github.com/Rain-kl/Wavelet/pkg/push"
@@ -23,8 +27,7 @@ import ("bytes"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/common/response")
)
var adminLoginEvent = EventMetadata{
Key: "admin_login",
@@ -187,7 +190,7 @@ func TestEventTrigger(t *testing.T) {
event.Enabled = true
event.Channels = []string{"mock_channel"}
event.Targets = []string{"admin_user"}
err = dbConn.Save(&event).Error
err = repository.SavePushEvent(context.Background(), &event)
require.NoError(t, err)
// Trigger
@@ -244,8 +247,8 @@ func TestEventTrigger(t *testing.T) {
event.Enabled = true
event.Channels = []string{"mock_channel"}
event.Targets = []string{"user.username"} // 动态目标
err = dbConn.Save(&event).Error
event.Targets = []string{"user.username"}
err = repository.SavePushEvent(context.Background(), &event)
require.NoError(t, err)
// Trigger with empty body (simulates cron scheduler triggering)
@@ -369,7 +372,7 @@ func TestPushRouters(t *testing.T) {
// 2. 为该事件关联渠道后,再切换开启,应当成功
event.Channels = []string{"email"}
dbConn.Save(&event)
_ = repository.SavePushEvent(context.Background(), &event)
req2, _ := http.NewRequest("POST", "/api/v1/admin/push/events/"+strconv.FormatUint(event.ID, 10)+"/toggle", nil)
w2 := httptest.NewRecorder()
+31 -3
View File
@@ -9,6 +9,8 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/db/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
@@ -35,7 +37,22 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error {
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
return repository.UpdateUserActive(ctx, id, active)
var tokens []model.AccessToken
if !active {
_ = db.DB(ctx).Where("user_id = ?", id).Find(&tokens).Error
}
err = repository.UpdateUserActive(ctx, id, active)
if err == nil {
oauth.InvalidateCachedUser(ctx, id)
if !active {
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
}
return err
}
func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
@@ -49,7 +66,18 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
if flags.IsAdmin {
return errors.New(cannotDelete)
}
return repository.DeleteUserWithRelations(ctx, targetID)
var tokens []model.AccessToken
_ = db.DB(ctx).Where("user_id = ?", targetID).Find(&tokens).Error
err = repository.DeleteUserWithRelations(ctx, targetID)
if err == nil {
oauth.InvalidateCachedUser(ctx, targetID)
for _, token := range tokens {
oauth.InvalidateCachedToken(ctx, token.TokenHash)
}
}
return err
}
func createUser(ctx context.Context, req createUserRequest) (model.User, error) {
@@ -103,4 +131,4 @@ func createUser(ctx context.Context, req createUserRequest) (model.User, error)
return model.User{}, err
}
return newUser, nil
}
}
+6 -3
View File
@@ -12,6 +12,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
@@ -103,7 +104,8 @@ func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbidde
return true
}
}
response.AbortInternal(c, msg)
logger.ErrorF(c.Request.Context(), "Admin user error: %v", err)
response.AbortInternal(c, "内部服务器错误")
return true
}
@@ -129,7 +131,8 @@ func ListUsers(c *gin.Context) {
total, modelUsers, err := listUsers(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err)
response.AbortInternal(c, "获取用户列表失败")
return
}
@@ -285,4 +288,4 @@ func CreateUser(c *gin.Context) {
}
c.JSON(http.StatusOK, response.OK(toUser(newUser)))
}
}
+20 -21
View File
@@ -6,9 +6,15 @@ package cap
import (
"net/http"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response"
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
)
// ChallengeResponse is a local type alias for the pkg/cap.ChallengeResponse struct
type ChallengeResponse = pkgcap.ChallengeResponse
type challengeRequest struct {
Scope string `json:"scope" form:"scope"`
}
@@ -26,8 +32,8 @@ type redeemRequest struct {
// @Accept json
// @Produce json
// @Param request body challengeRequest false "可选范围限制参数"
// @Success 200 {object} cap.ChallengeResponse "成功返回 PoW 难题"
// @Failure 500 {object} RedeemResponse "内部服务错误"
// @Success 200 {object} response.Any{data=cap.ChallengeResponse} "成功返回 PoW 难题"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/cap/challenge [post]
func Challenge(c *gin.Context) {
var req challengeRequest
@@ -40,14 +46,12 @@ func Challenge(c *gin.Context) {
mgr := GetDefaultManager()
resp, err := mgr.Generate(c.Request.Context(), req.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, RedeemResponse{
Success: false,
Error: err.Error(),
})
logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err)
response.AbortInternal(c, "生成验证难题失败,请稍后再试")
return
}
c.JSON(http.StatusOK, resp)
c.JSON(http.StatusOK, response.OK(resp))
}
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
@@ -57,17 +61,14 @@ func Challenge(c *gin.Context) {
// @Accept json
// @Produce json
// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组"
// @Success 200 {object} RedeemResponse "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} RedeemResponse "参数错误或核销失败"
// @Failure 500 {object} RedeemResponse "内部服务错误"
// @Success 200 {object} response.Any{data=cap.RedeemResponse} "核销成功,返回 X-Cap-Token"
// @Failure 400 {object} response.Any "参数错误或核销失败"
// @Failure 500 {object} response.Any "内部服务错误"
// @Router /api/cap/redeem [post]
func Redeem(c *gin.Context) {
var req redeemRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, RedeemResponse{
Success: false,
Error: "无效的参数",
})
response.AbortBadRequest(c, "无效的参数")
return
}
@@ -78,17 +79,15 @@ func Redeem(c *gin.Context) {
mgr := GetDefaultManager()
resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, RedeemResponse{
Success: false,
Error: err.Error(),
})
logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err)
response.AbortInternal(c, "校验验证解答失败,请稍后再试")
return
}
if !resp.Success {
c.JSON(http.StatusBadRequest, resp)
response.AbortBadRequest(c, resp.Error)
return
}
c.JSON(http.StatusOK, resp)
c.JSON(http.StatusOK, response.OK(resp))
}
+12 -4
View File
@@ -46,10 +46,14 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
t.Fatalf("expected 200 OK, got %d. Body: %s", w.Code, w.Body.String())
}
var challengeResp pkgcap.ChallengeResponse
if err := json.Unmarshal(w.Body.Bytes(), &challengeResp); err != nil {
var envelope struct {
ErrorMsg string `json:"error_msg"`
Data pkgcap.ChallengeResponse `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &envelope); err != nil {
t.Fatalf("failed to unmarshal challenge response: %v", err)
}
challengeResp := envelope.Data
if challengeResp.Token == "" {
t.Fatalf("expected token in challenge response")
@@ -99,10 +103,14 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
t.Fatalf("expected 200 OK for redeem, got %d. Body: %s", w.Code, w.Body.String())
}
var redeemResp RedeemResponse
if err := json.Unmarshal(w.Body.Bytes(), &redeemResp); err != nil {
var redeemEnvelope struct {
ErrorMsg string `json:"error_msg"`
Data RedeemResponse `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &redeemEnvelope); err != nil {
t.Fatalf("failed to unmarshal redeem response: %v", err)
}
redeemResp := redeemEnvelope.Data
if !redeemResp.Success || redeemResp.Token == "" {
t.Fatalf("redeem failed or returned empty token: %+v", redeemResp)
+152
View File
@@ -0,0 +1,152 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package oauth
import (
"context"
"fmt"
"sync"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
type cacheEntry struct {
value any
expiredAt time.Time
}
type memoryCache struct {
sync.RWMutex
items map[string]cacheEntry
}
var localCache = &memoryCache{
items: make(map[string]cacheEntry),
}
func (c *memoryCache) Set(key string, val any, ttl time.Duration) {
c.Lock()
defer c.Unlock()
c.items[key] = cacheEntry{
value: val,
expiredAt: time.Now().Add(ttl),
}
}
func (c *memoryCache) Get(key string) (any, bool) {
c.RLock()
item, ok := c.items[key]
if !ok {
c.RUnlock()
return nil, false
}
if time.Now().After(item.expiredAt) {
c.RUnlock()
c.Lock()
if item, ok = c.items[key]; ok && time.Now().After(item.expiredAt) {
delete(c.items, key)
}
c.Unlock()
return nil, false
}
c.RUnlock()
return item.value, true
}
func (c *memoryCache) Delete(key string) {
c.Lock()
defer c.Unlock()
delete(c.items, key)
}
const (
tokenCacheTTL = 5 * time.Minute
userCacheTTL = 5 * time.Minute
)
func tokenCacheKey(tokenHash string) string {
return "oauth:token:" + tokenHash
}
func userCacheKey(userID uint64) string {
return fmt.Sprintf("oauth:user:%d", userID)
}
// GetCachedToken 获取缓存的 AccessToken
func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) {
key := tokenCacheKey(tokenHash)
if val, ok := localCache.Get(key); ok {
if token, ok := val.(*model.AccessToken); ok {
return token, nil
}
}
if db.Redis != nil {
var token model.AccessToken
if err := db.GetJSON(ctx, key, &token); err == nil {
// Write back to local cache
localCache.Set(key, &token, tokenCacheTTL)
return &token, nil
}
}
return nil, fmt.Errorf("cache miss")
}
// SetCachedToken 设置 AccessToken 缓存
func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) {
key := tokenCacheKey(tokenHash)
localCache.Set(key, token, tokenCacheTTL)
if db.Redis != nil {
_ = db.SetJSON(ctx, key, token, tokenCacheTTL)
}
}
// InvalidateCachedToken 吊销/删除 token 缓存
func InvalidateCachedToken(ctx context.Context, tokenHash string) {
key := tokenCacheKey(tokenHash)
localCache.Delete(key)
if db.Redis != nil {
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
}
}
// GetCachedUser 获取缓存的 User
func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) {
key := userCacheKey(userID)
if val, ok := localCache.Get(key); ok {
if u, ok := val.(*model.User); ok {
return u, nil
}
}
if db.Redis != nil {
var u model.User
if err := db.GetJSON(ctx, key, &u); err == nil {
// Write back to local cache
localCache.Set(key, &u, userCacheTTL)
return &u, nil
}
}
return nil, fmt.Errorf("cache miss")
}
// SetCachedUser 设置 User 缓存
func SetCachedUser(ctx context.Context, userID uint64, u *model.User) {
key := userCacheKey(userID)
localCache.Set(key, u, userCacheTTL)
if db.Redis != nil {
_ = db.SetJSON(ctx, key, u, userCacheTTL)
}
}
// InvalidateCachedUser 吊销/失效 User 缓存
func InvalidateCachedUser(ctx context.Context, userID uint64) {
key := userCacheKey(userID)
localCache.Delete(key)
if db.Redis != nil {
_ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err()
}
}
+29 -13
View File
@@ -31,15 +31,26 @@ type loginRequiredAuditLog struct {
func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) {
tokenHash := model.HashToken(tokenStr)
var tokenRecord model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err != nil {
return nil, nil, err
tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil {
var dbToken model.AccessToken
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&dbToken).Error; err != nil {
return nil, nil, err
}
tokenRecord = &dbToken
SetCachedToken(ctx, tokenHash, tokenRecord)
}
var user model.User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err != nil {
return nil, nil, err
user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive {
var dbUser model.User
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil {
return nil, nil, err
}
user = &dbUser
SetCachedUser(ctx, tokenRecord.UserID, user)
}
return &user, &tokenRecord, nil
return user, tokenRecord, nil
}
// GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error
@@ -74,11 +85,16 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
return nil, errors.New("unauthorized")
}
var user model.User
// 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
user, err := GetCachedUser(ctx, userID)
if err != nil || !user.IsActive {
var dbUser model.User
// load user from db to make sure is active
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&dbUser)
if tx.Error != nil {
return nil, tx.Error
}
user = &dbUser
SetCachedUser(ctx, userID, user)
}
// 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致
@@ -99,7 +115,7 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
return nil, errors.New("system user is not allowed to login")
}
return &user, nil
return user, nil
}
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
+11 -3
View File
@@ -4,15 +4,16 @@
package oauth
import ("net/http"
import (
"net/http"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/common/response")
"github.com/Rain-kl/Wavelet/internal/common/response"
)
// BasicUserInfo 用户基本信息结构体
type BasicUserInfo struct {
@@ -94,6 +95,13 @@ func Logout(c *gin.Context) {
username := session.Get(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 {
InvalidateCachedUser(c.Request.Context(), id)
} else if idFloat, ok := userID.(float64); ok {
InvalidateCachedUser(c.Request.Context(), uint64(idFloat))
} else if idInt, ok := userID.(int); ok && idInt >= 0 {
InvalidateCachedUser(c.Request.Context(), uint64(idInt))
}
}
session.Options(GetSessionOptions(-1))
session.Clear()
+7 -3
View File
@@ -7,6 +7,7 @@ package handler
import (
"archive/zip"
"bufio"
"bytes"
"crypto/sha256"
"encoding/hex"
@@ -247,8 +248,12 @@ func BatchDownloadFiles(c *gin.Context) {
c.Header("Content-Type", "application/zip")
c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"")
zipWriter := zip.NewWriter(c.Writer)
defer func() { _ = zipWriter.Close() }()
bufferedWriter := bufio.NewWriter(c.Writer)
zipWriter := zip.NewWriter(bufferedWriter)
defer func() {
_ = zipWriter.Close()
_ = bufferedWriter.Flush()
}()
usedNames := make(map[string]int)
@@ -326,4 +331,3 @@ func detectMimeType(buf *bytes.Buffer, header *multipart.FileHeader, size int64)
}
return mimeType
}
+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
+8 -1
View File
@@ -15,6 +15,7 @@ import (
"github.com/ClickHouse/clickhouse-go/v2"
"github.com/Rain-kl/Wavelet/internal/config"
"github.com/pressly/goose/v3"
"github.com/pressly/goose/v3/database"
)
const (
@@ -63,11 +64,17 @@ func MigrateClickHouse() {
log.Fatalf("[ClickHouse] get sub fs failed: %v\n", err)
}
store, err := database.NewStore(database.DialectClickHouse, clickhouseGooseVersionTable)
if err != nil {
closeClickHouseDB(sqlDB)
log.Fatalf("[ClickHouse] create goose store failed: %v\n", err)
}
provider, err := goose.NewProvider(
"clickhouse",
sqlDB,
subFS,
goose.WithTableName(clickhouseGooseVersionTable),
goose.WithStore(store),
goose.WithDisableGlobalRegistry(true),
)
if err != nil {
+22 -4
View File
@@ -30,8 +30,10 @@ type systemConfigInvalidationMessage struct {
}
var (
systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize})
systemConfigListenerOnce sync.Once
systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize})
systemConfigListenerOnce sync.Once
systemConfigListenerCtx context.Context
systemConfigListenerCancel context.CancelFunc
)
func ensureSystemConfigCacheListener() {
@@ -43,12 +45,19 @@ func startSystemConfigCacheInvalidationListener() {
return
}
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
go func() {
pubsub := db.Redis.Subscribe(context.Background(), SystemConfigInvalidationChannel)
pubsub := db.Redis.Subscribe(systemConfigListenerCtx, SystemConfigInvalidationChannel)
defer func() {
_ = pubsub.Close()
}()
go func() {
<-systemConfigListenerCtx.Done()
_ = pubsub.Close()
}()
for msg := range pubsub.Channel() {
var payload systemConfigInvalidationMessage
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
@@ -64,6 +73,15 @@ func startSystemConfigCacheInvalidationListener() {
}()
}
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
func StopSystemConfigCacheListener() {
if systemConfigListenerCancel != nil {
systemConfigListenerCancel()
systemConfigListenerCancel = nil
}
systemConfigListenerOnce = sync.Once{}
}
func cloneSystemConfig(sc model.SystemConfig) model.SystemConfig {
return sc
}
@@ -117,4 +135,4 @@ func InvalidateAllSystemConfigCaches(ctx context.Context) error {
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
systemConfigRAMCache.InvalidateAll()
}
}