mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 23:16:37 +08:00
refactor(core): optimize performance, fix concurrency and clean up AGENTS.md design violations
- Concurrency: Added lock protection to WebSocket writes, fixed timer leaks, and prevented config cache listener context leaks. - Performance: Added memory cache in ObservabilityBufferStore, periodic cleaning in CH Deduplicator, and buffered ZIP batch download writes. - Design: Introduced Redis caching for OAuth session/tokens, sanitized raw DB error messages, segregated handlers and logics, and standard CAP response envelopes.
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -26,8 +26,10 @@ type ObservabilityBufferRecord struct {
|
||||
|
||||
// ObservabilityBufferStore persists observability records to disk for replay on heartbeat.
|
||||
type ObservabilityBufferStore struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
path string
|
||||
mu sync.Mutex
|
||||
cache []ObservabilityBufferRecord
|
||||
cacheLoaded bool
|
||||
}
|
||||
|
||||
// NewObservabilityBufferStore creates a store backed by the file at path.
|
||||
@@ -174,21 +176,34 @@ func (s *ObservabilityBufferStore) Ack(windowStartedAtUnix []int64, retainAfterU
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) loadUnlocked() ([]ObservabilityBufferRecord, error) {
|
||||
if s.cacheLoaded {
|
||||
copied := make([]ObservabilityBufferRecord, len(s.cache))
|
||||
copy(copied, s.cache)
|
||||
return copied, nil
|
||||
}
|
||||
data, err := os.ReadFile(s.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
s.cache = []ObservabilityBufferRecord{}
|
||||
s.cacheLoaded = true
|
||||
return []ObservabilityBufferRecord{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if len(data) == 0 {
|
||||
s.cache = []ObservabilityBufferRecord{}
|
||||
s.cacheLoaded = true
|
||||
return []ObservabilityBufferRecord{}, nil
|
||||
}
|
||||
var records []ObservabilityBufferRecord
|
||||
if err = json.Unmarshal(data, &records); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return records, nil
|
||||
s.cache = records
|
||||
s.cacheLoaded = true
|
||||
copied := make([]ObservabilityBufferRecord, len(s.cache))
|
||||
copy(copied, s.cache)
|
||||
return copied, nil
|
||||
}
|
||||
|
||||
func (s *ObservabilityBufferStore) saveUnlocked(records []ObservabilityBufferRecord) error {
|
||||
@@ -199,7 +214,12 @@ func (s *ObservabilityBufferStore) saveUnlocked(records []ObservabilityBufferRec
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(s.path, data, stateFilePerm)
|
||||
if err := os.WriteFile(s.path, data, stateFilePerm); err != nil {
|
||||
return err
|
||||
}
|
||||
s.cache = records
|
||||
s.cacheLoaded = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// ObservabilityWindowStartedAt calculates the start of the 60-second window for the given metrics, openresty observation, or traffic report.
|
||||
|
||||
@@ -241,14 +241,49 @@ func switchPagesCurrentDir(baseDir string, deploymentID uint, releaseDir string)
|
||||
if err := os.MkdirAll(filepath.Dir(currentDir), pagesDirPerm); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := os.Stat(currentDir); err == nil {
|
||||
|
||||
relTarget, err := filepath.Rel(filepath.Dir(currentDir), releaseDir)
|
||||
if err != nil {
|
||||
relTarget = releaseDir
|
||||
}
|
||||
|
||||
// Try creating a temporary symlink first to check if symlinks are supported/feasible
|
||||
tmpSymlink := currentDir + ".tmp"
|
||||
_ = os.Remove(tmpSymlink)
|
||||
|
||||
symlinkErr := os.Symlink(relTarget, tmpSymlink)
|
||||
if symlinkErr != nil {
|
||||
return fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir)
|
||||
}
|
||||
|
||||
// Symlink is supported, proceed with symlink swap
|
||||
_ = os.Remove(tmpSymlink)
|
||||
|
||||
if _, err := os.Lstat(currentDir); err == nil {
|
||||
if err := os.Rename(currentDir, previousDir); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err := os.Symlink(relTarget, currentDir); err != nil {
|
||||
if _, restoreErr := os.Lstat(previousDir); restoreErr == nil {
|
||||
_ = os.Rename(previousDir, currentDir)
|
||||
}
|
||||
return err
|
||||
}
|
||||
_ = os.RemoveAll(previousDir)
|
||||
return nil
|
||||
}
|
||||
|
||||
func fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir string) error {
|
||||
if _, err := os.Lstat(currentDir); err == nil {
|
||||
if err := os.Rename(currentDir, previousDir); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := copyPagesDir(releaseDir, currentDir); err != nil {
|
||||
_ = os.RemoveAll(currentDir)
|
||||
if _, restoreErr := os.Stat(previousDir); restoreErr == nil {
|
||||
if _, restoreErr := os.Lstat(previousDir); restoreErr == nil {
|
||||
_ = os.Rename(previousDir, currentDir)
|
||||
}
|
||||
return err
|
||||
|
||||
@@ -6,9 +6,14 @@ package cap
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||
pkgcap "github.com/Rain-kl/Wavelet/pkg/cap"
|
||||
"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 +31,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 +45,11 @@ 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(),
|
||||
})
|
||||
response.AbortInternal(c, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, resp)
|
||||
c.JSON(http.StatusOK, response.OK(resp))
|
||||
}
|
||||
|
||||
// Redeem 提交 PoW 解答并兑换一次性凭证 Token
|
||||
@@ -57,17 +59,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 +77,14 @@ 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(),
|
||||
})
|
||||
response.AbortInternal(c, err.Error())
|
||||
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))
|
||||
}
|
||||
|
||||
@@ -63,8 +63,10 @@ func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig,
|
||||
|
||||
// SleepContext pauses execution for the given duration or until the context is canceled.
|
||||
func SleepContext(ctx context.Context, d time.Duration) {
|
||||
t := time.NewTimer(d)
|
||||
defer t.Stop()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-time.After(d):
|
||||
case <-t.C:
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/apps/flared/config"
|
||||
@@ -225,15 +227,18 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
backoff = 1 * time.Second
|
||||
}
|
||||
|
||||
t := time.NewTimer(backoff)
|
||||
select {
|
||||
case <-procCtx.Done():
|
||||
t.Stop()
|
||||
return
|
||||
case <-time.After(backoff):
|
||||
case <-t.C:
|
||||
backoff *= 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
}
|
||||
t.Stop()
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -350,10 +355,13 @@ func ensureNoOrphanProcess(pidPath string) {
|
||||
}
|
||||
process, err := os.FindProcess(pid)
|
||||
if err == nil && process != nil {
|
||||
slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath)
|
||||
_ = process.Kill()
|
||||
// Wait a little bit to ensure the OS has reclaimed ports
|
||||
time.Sleep(orphanProcessKillDelay)
|
||||
err = process.Signal(syscall.Signal(0))
|
||||
if err == nil || errors.Is(err, os.ErrPermission) {
|
||||
slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath)
|
||||
_ = process.Kill()
|
||||
// Wait a little bit to ensure the OS has reclaimed ports
|
||||
time.Sleep(orphanProcessKillDelay)
|
||||
}
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
// 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()
|
||||
defer c.RUnlock()
|
||||
item, ok := c.items[key]
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
if time.Now().After(item.expiredAt) {
|
||||
return nil, false
|
||||
}
|
||||
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
|
||||
|
||||
@@ -95,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()
|
||||
|
||||
@@ -11,12 +11,16 @@ import (
|
||||
const dedupTTL = 2 * time.Minute
|
||||
|
||||
type dedupSet struct {
|
||||
mu sync.Mutex
|
||||
keys map[string]time.Time
|
||||
mu sync.Mutex
|
||||
keys map[string]time.Time
|
||||
lastCleanup time.Time
|
||||
}
|
||||
|
||||
func newDedupSet() *dedupSet {
|
||||
return &dedupSet{keys: make(map[string]time.Time)}
|
||||
return &dedupSet{
|
||||
keys: make(map[string]time.Time),
|
||||
lastCleanup: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// markIfNew records key when it has not been seen within dedupTTL.
|
||||
@@ -29,11 +33,16 @@ func (s *dedupSet) markIfNew(key string) bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
for existing, expiresAt := range s.keys {
|
||||
if now.After(expiresAt) {
|
||||
delete(s.keys, existing)
|
||||
// Periodically clean up all expired keys (e.g., every 30 seconds)
|
||||
if now.Sub(s.lastCleanup) >= 30*time.Second {
|
||||
for existing, expiresAt := range s.keys {
|
||||
if now.After(expiresAt) {
|
||||
delete(s.keys, existing)
|
||||
}
|
||||
}
|
||||
s.lastCleanup = now
|
||||
}
|
||||
|
||||
if expiresAt, exists := s.keys[key]; exists && now.Before(expiresAt) {
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ package frps
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
service "github.com/Rain-kl/Wavelet/pkg/protocol"
|
||||
@@ -322,10 +324,13 @@ func ensureNoOrphanProcess(pidPath string) {
|
||||
}
|
||||
process, err := os.FindProcess(pid)
|
||||
if err == nil && process != nil {
|
||||
slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath)
|
||||
_ = process.Kill()
|
||||
// Wait a little bit to ensure the OS has reclaimed ports
|
||||
time.Sleep(frpsOrphanProcessCleanupDelay)
|
||||
err = process.Signal(syscall.Signal(0))
|
||||
if err == nil || errors.Is(err, os.ErrPermission) {
|
||||
slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath)
|
||||
_ = process.Kill()
|
||||
// Wait a little bit to ensure the OS has reclaimed ports
|
||||
time.Sleep(frpsOrphanProcessCleanupDelay)
|
||||
}
|
||||
}
|
||||
_ = os.Remove(pidPath)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,125 @@ 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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user