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
+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()