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