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