refactor(architecture): eliminate internal package and complete cordis single-owner model and repository migration

- Physically purged all legacy internal/ packages, centralized pkg/model/ and pkg/repository/
- Migrated domain models and database repositories into self-contained owner plugins (user, auth, message_gateway, admin, upload, risk_control)
- Decoupled cross-plugin interactions via pure core/contracts and typed EventBus
- Ensured 100% test coverage pass, zero data races (-race clean), and 0 lint issues in make code-check
This commit is contained in:
ryan
2026-08-28 08:40:43 +08:00
parent 1f348fd425
commit fb6a3edb89
323 changed files with 8222 additions and 17693 deletions
+2 -2
View File
@@ -7,8 +7,8 @@ import (
"net/http"
"strconv"
persistence "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/shared/response"
persistence "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-gonic/gin"
)
+6 -8
View File
@@ -10,10 +10,8 @@ import (
"github.com/gin-gonic/gin"
"github.com/Rain-kl/Wavelet/internal/infra/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/diskcache"
)
type updateCacheConfigRequest struct {
@@ -61,17 +59,17 @@ func UpdateCacheConfig(c *gin.Context) {
ctx := c.Request.Context()
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheMaxSizeMB, strconv.FormatInt(req.MaxSizeMB, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheTTLMinutes, strconv.FormatInt(req.TTLMinutes, 10)); err != nil {
response.AbortInternal(c, err.Error())
return
}
if err := saveOrUpdateCacheConfig(ctx, model.ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
if err := saveOrUpdateCacheConfig(ctx, ConfigKeyDiskCacheLRUEnabled, strconv.FormatBool(req.LRUEnabled)); err != nil {
response.AbortInternal(c, err.Error())
return
}
@@ -101,5 +99,5 @@ func ClearCache(c *gin.Context) {
}
func saveOrUpdateCacheConfig(ctx context.Context, key, value string) error {
return repository.SaveOrUpdateSystemConfig(ctx, key, value)
return SaveOrUpdateSystemConfig(ctx, key, value)
}
+36 -37
View File
@@ -13,15 +13,12 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
mail "github.com/Rain-kl/Wavelet/pkg/mail"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/plugins/domain/cap"
"github.com/Rain-kl/Wavelet/plugins/domain/upload"
"github.com/Rain-kl/Wavelet/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -54,7 +51,7 @@ type UpdateSystemConfigRequest struct {
// @Router /api/v1/config/public [get]
func GetPublicConfig(c *gin.Context) {
ctx := c.Request.Context()
configs, err := repository.ListVisibleSystemConfigs(ctx)
configs, err := ListVisibleSystemConfigs(ctx)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -77,7 +74,7 @@ func GetPublicConfig(c *gin.Context) {
// @Router /robots.txt [get]
func GetRobotsTXT(c *gin.Context) {
ctx := c.Request.Context()
enabled, err := repository.GetBoolByKey(ctx, model.ConfigKeySearchEngineIndexingEnabled)
enabled, err := GetBoolByKey(ctx, ConfigKeySearchEngineIndexingEnabled)
content := "User-Agent: *\nDisallow: /\n"
if err == nil && enabled {
content = "User-Agent: *\nAllow: /\n"
@@ -129,7 +126,7 @@ func CreateSystemConfig(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param type query string false "配置类型(system/business)"
// @Success 200 {object} response.Any{data=[]model.SystemConfig} "系统配置列表"
// @Success 200 {object} response.Any{data=[]SystemConfig} "系统配置列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
@@ -155,7 +152,7 @@ func ListSystemConfigs(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param key path string true "配置键"
// @Success 200 {object} response.Any{data=model.SystemConfig} "系统配置详情"
// @Success 200 {object} response.Any{data=SystemConfig} "系统配置详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "配置不存在"
@@ -222,14 +219,14 @@ func UpdateSystemConfig(c *gin.Context) {
}
func isProtectedConfigKey(key string) bool {
return key == model.ConfigKeyLogDatabase || key == model.ConfigKeyLogDBMigration
return key == ConfigKeyLogDatabase || key == ConfigKeyLogDBMigration
}
func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) error {
if isProtectedConfigKey(req.Key) {
return errors.New(protectedConfigKeyMessage)
}
exists, err := repository.SystemConfigExists(ctx, req.Key)
exists, err := SystemConfigExists(ctx, req.Key)
if err != nil {
return err
}
@@ -237,43 +234,43 @@ func createSystemConfig(ctx context.Context, req CreateSystemConfigRequest) erro
return errors.New(ConfigKeyExists)
}
config := model.SystemConfig{
config := SystemConfig{
Key: req.Key,
Value: req.Value,
Type: req.Type,
Visibility: req.Visibility,
Description: req.Description,
}
if err := repository.CreateSystemConfig(ctx, &config); err != nil {
if err := CreateSystemConfigRecord(ctx, &config); err != nil {
return err
}
invalidateSystemConfigCaches(ctx, req.Key)
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
return nil
}
func listSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) {
return repository.ListAdminSystemConfigs(ctx, configType)
func listSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
return ListAdminSystemConfigs(ctx, configType)
}
func getSystemConfig(ctx context.Context, key string) (model.SystemConfig, error) {
return repository.GetAdminSystemConfigByKey(ctx, key)
func getSystemConfig(ctx context.Context, key string) (SystemConfig, error) {
return GetAdminSystemConfigByKey(ctx, key)
}
func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigRequest) error {
if isProtectedConfigKey(key) {
return errors.New(protectedConfigKeyMessage)
}
config, err := repository.GetAdminSystemConfigByKey(ctx, key)
config, err := GetAdminSystemConfigByKey(ctx, key)
if err != nil {
return err
}
var originalDriver objectstore.Driver
if key == model.ConfigKeyStorageConfig {
if key == ConfigKeyStorageConfig {
var currentCfg objectstore.Config
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
@@ -294,7 +291,7 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
updates["visibility"] = *req.Visibility
config.Visibility = *req.Visibility
}
if key != model.ConfigKeySMTPPassword || req.Value != maskedConfigValue {
if key != ConfigKeySMTPPassword || req.Value != maskedConfigValue {
updates["value"] = req.Value
config.Value = req.Value
}
@@ -318,7 +315,7 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
originalDriver objectstore.Driver,
newValue string,
) {
if key != model.ConfigKeyStorageConfig || originalDriver == "" {
if key != ConfigKeyStorageConfig || originalDriver == "" {
return
}
@@ -330,7 +327,7 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
return
}
if err := repository.MarkFailedTaskExecutionsSucceededTx(
if err := MarkFailedTaskExecutionsSucceededTx(
tx,
"storage:migrate",
"存储配置直接更新,故障迁移任务自动标记为已解决",
@@ -341,7 +338,7 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
}
func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := repository.InvalidateSystemConfigCache(ctx, key); err != nil {
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if cap.IsRuntimeConfigKey(key) {
@@ -352,18 +349,20 @@ func invalidateSystemConfigCaches(ctx context.Context, key string) {
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if key == model.ConfigKeyStorageConfig {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
if key == ConfigKeyStorageConfig {
if db.Redis != nil {
_ = db.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
}
objectstore.ResetCache()
objectstore.PublishCacheInvalidation(ctx)
}
if key == model.ConfigKeyFileAccessWhitelist {
upload.ResetAccessCaches()
upload.PublishAccessCacheInvalidation(ctx)
if key == ConfigKeyFileAccessWhitelist {
if db.Redis != nil {
_ = db.Redis.Publish(ctx, "upload:access_cache:invalidate", "reset").Err()
}
}
if err := repository.InvalidateVisibleSystemConfigsCache(ctx); err != nil {
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
}
@@ -404,7 +403,7 @@ func TestSMTP(c *gin.Context) {
password := req.SMTPPassword
if password == maskedConfigValue {
if sc, err := repository.GetSystemConfigByKey(c.Request.Context(), model.ConfigKeySMTPPassword); err == nil {
if sc, err := GetSystemConfigByKey(c.Request.Context(), ConfigKeySMTPPassword); err == nil {
password = sc.Value
}
}
@@ -449,9 +448,9 @@ func maskSensitiveConfig(key, value string) string {
return value
}
switch key {
case model.ConfigKeySMTPPassword:
case ConfigKeySMTPPassword:
return maskedConfigValue
case model.ConfigKeyStorageConfig:
case ConfigKeyStorageConfig:
var cfg objectstore.Config
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
masked := objectstore.MaskSecrets(cfg)
@@ -494,8 +493,8 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
if err := db.DB(ctx).Model(&model.Upload{}).
Where("status != ?", model.UploadStatusDeleted).
if err := db.DB(ctx).Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
+3 -3
View File
@@ -19,9 +19,9 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/response"
)
const (
+25 -15
View File
@@ -15,13 +15,12 @@ import (
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/risk_control"
"github.com/gin-gonic/gin"
@@ -164,8 +163,10 @@ func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessL
username := c.Query("username")
if username != "" {
userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username)
if err != nil {
var userIDs []uint64
if err := db.DB(ctx).Table("w_users").
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
Pluck("id", &userIDs).Error; err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err)
}
filter.UserIDs = userIDs
@@ -213,7 +214,12 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
}
userMap := make(map[uint64]struct{ Username, Nickname string })
if users, err := repository.ListUsersByIDs(ctx, userIDs); err == nil {
var users []struct {
ID uint64
Username string
Nickname string
}
if err := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
}
@@ -402,8 +408,12 @@ func GetLogsAnalytics(c *gin.Context) {
Username string
Nickname string
})
users, errProfile := repository.ListUsersByIDs(ctx, userIDs)
if errProfile == nil {
var users []struct {
ID uint64
Username string
Nickname string
}
if errProfile := db.DB(ctx).Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
for _, u := range users {
userProfileMap[u.ID] = struct {
Username string
@@ -445,7 +455,7 @@ func getUpgrader() *websocket.Upgrader {
// 2. 检查配置的允许跨域 Origin (Check allowed origins in system config)
ctx := r.Context()
if sc, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
if sc, err := GetSystemConfigByKey(ctx, ConfigKeyServerAddress); err == nil && sc.Value != "" {
originToCheck := strings.TrimRight(strings.TrimSpace(origin), "/")
allowedOrigins := strings.Split(sc.Value, ",")
for _, allowed := range allowedOrigins {
@@ -632,7 +642,7 @@ func validateSwitch(ctx context.Context, target string) error {
}
func currentLogDatabase(ctx context.Context) (string, error) {
cfg, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyLogDatabase)
cfg, err := GetSystemConfigByKey(ctx, ConfigKeyLogDatabase)
if err != nil {
return "", fmt.Errorf("读取日志主库失败: %w", err)
}
@@ -643,11 +653,11 @@ func currentLogDatabase(ctx context.Context) (string, error) {
}
func setMigrationFlag(ctx context.Context, v string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDBMigration, v)
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDBMigration, v)
}
func flipLogDatabase(ctx context.Context, target string) error {
return repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, target)
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
}
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
+7 -9
View File
@@ -16,12 +16,10 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/repository/logstore"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/config"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/persistence/logstore"
"github.com/Rain-kl/Wavelet/pkg/response"
)
var startTime = time.Now()
@@ -200,16 +198,16 @@ func GetLogDatabaseStatus(c *gin.Context) {
ActiveDatabase: activeDB,
Migration: migration,
RetentionDays: map[string]int{
logDBNamePostgres: retentionOr(ctx, model.ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, model.ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, model.ConfigKeyLogRetentionDaysClickHouse),
logDBNamePostgres: retentionOr(ctx, ConfigKeyLogRetentionDaysPostgres),
logDBNameSQLite: retentionOr(ctx, ConfigKeyLogRetentionDaysSQLite),
logDBNameClickHouse: retentionOr(ctx, ConfigKeyLogRetentionDaysClickHouse),
},
AvailableTargets: availableLogTargets(activeDB),
}))
}
func retentionOr(ctx context.Context, key string) int {
v, err := repository.GetIntByKey(ctx, key)
v, err := GetIntByKey(ctx, key)
if err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
logger.ErrorF(ctx, "读取日志保留天数配置失败 key=%s: %v", key, err)
+16 -18
View File
@@ -14,12 +14,10 @@ import (
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
"github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/infra/task/scheduler"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/task"
"github.com/Rain-kl/Wavelet/pkg/task/scheduler"
)
// ListTaskTypes 获取支持的任务类型列表
@@ -107,7 +105,7 @@ func DispatchTask(c *gin.Context) {
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/executions [get]
func ListTaskExecutions(c *gin.Context) {
var req model.ListTaskExecutionsRequest
var req ListTaskExecutionsRequest
if err := c.ShouldBindQuery(&req); err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -119,7 +117,7 @@ func ListTaskExecutions(c *gin.Context) {
}
}
executions, total, err := repository.ListTaskExecutions(c.Request.Context(), req)
executions, total, err := ListTaskExecutionRecords(c.Request.Context(), req)
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -140,7 +138,7 @@ func ListTaskExecutions(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param id path int true "任务执行记录 ID"
// @Success 200 {object} response.Any{data=model.TaskExecution} "任务执行详情"
// @Success 200 {object} response.Any{data=TaskExecution} "任务执行详情"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -153,7 +151,7 @@ func GetTaskExecution(c *gin.Context) {
return
}
execution, err := repository.GetTaskExecutionByID(c.Request.Context(), id)
execution, err := GetTaskExecutionByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, TaskNotFound)
return
@@ -206,12 +204,12 @@ func RetryTask(c *gin.Context) {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Schedule} "定时任务列表"
// @Success 200 {object} response.Any{data=[]Schedule} "定时任务列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) {
schedules, err := repository.ListSchedules(c.Request.Context())
schedules, err := ListSchedulesRecord(c.Request.Context())
if err != nil {
response.AbortInternal(c, err.Error())
return
@@ -236,7 +234,7 @@ type CreateScheduleRequest struct {
// @Produce json
// @Security SessionCookie
// @Param request body CreateScheduleRequest true "创建定时任务请求参数"
// @Success 200 {object} response.Any{data=model.Schedule} "创建成功的定时任务信息"
// @Success 200 {object} response.Any{data=Schedule} "创建成功的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、异步任务类型不存在或参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -273,7 +271,7 @@ func CreateSchedule(c *gin.Context) {
return
}
schedule := &model.Schedule{
schedule := &Schedule{
Name: req.Name,
TaskType: req.TaskType,
Cron: req.Cron,
@@ -281,7 +279,7 @@ func CreateSchedule(c *gin.Context) {
IsActive: *req.IsActive,
}
if err := repository.CreateSchedule(c.Request.Context(), schedule); err != nil {
if err := CreateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -312,7 +310,7 @@ type UpdateScheduleRequest struct {
// @Security SessionCookie
// @Param id path int true "定时任务 ID"
// @Param request body UpdateScheduleRequest true "修改定时任务请求参数"
// @Success 200 {object} response.Any{data=model.Schedule} "修改后的定时任务信息"
// @Success 200 {object} response.Any{data=Schedule} "修改后的定时任务信息"
// @Failure 400 {object} response.Any "Cron 表达式无效、参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -332,7 +330,7 @@ func UpdateSchedule(c *gin.Context) {
return
}
schedule, err := repository.GetScheduleByID(c.Request.Context(), id)
schedule, err := GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
return
@@ -368,7 +366,7 @@ func UpdateSchedule(c *gin.Context) {
schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive
if err := repository.UpdateSchedule(c.Request.Context(), schedule); err != nil {
if err := UpdateScheduleRecord(c.Request.Context(), schedule); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return
}
@@ -401,7 +399,7 @@ func DeleteSchedule(c *gin.Context) {
return
}
if err := repository.DeleteSchedule(c.Request.Context(), id); err != nil {
if err := DeleteScheduleRecord(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return
}
+24 -26
View File
@@ -11,9 +11,7 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/response"
)
// CreateTemplateRequest 创建模板请求
@@ -88,7 +86,7 @@ func CreateTemplate(c *gin.Context) {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]model.Template} "模板列表"
// @Success 200 {object} response.Any{data=[]Template} "模板列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
@@ -110,7 +108,7 @@ func ListTemplates(c *gin.Context) {
// @Produce json
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Success 200 {object} response.Any{data=model.Template} "模板详情"
// @Success 200 {object} response.Any{data=Template} "模板详情"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "模板不存在"
@@ -134,7 +132,7 @@ func GetTemplate(c *gin.Context) {
// @Security SessionCookie
// @Param key path string true "模板标识符"
// @Param request body UpdateTemplateRequest true "更新请求参数"
// @Success 200 {object} response.Any{data=model.Template} "更新成功"
// @Success 200 {object} response.Any{data=Template} "更新成功"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
@@ -178,16 +176,16 @@ func DeleteTemplate(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Template, error) {
exists, err := repository.TemplateExistsByKey(ctx, req.Key)
func createTemplate(ctx context.Context, req CreateTemplateRequest) (Template, error) {
exists, err := TemplateExistsByKey(ctx, req.Key)
if err != nil {
return model.Template{}, err
return Template{}, err
}
if exists {
return model.Template{}, errors.New(TemplateKeyExists)
return Template{}, errors.New(TemplateKeyExists)
}
tmpl := model.Template{
tmpl := Template{
Key: req.Key,
Name: req.Name,
Type: req.Type,
@@ -197,26 +195,26 @@ func createTemplate(ctx context.Context, req CreateTemplateRequest) (model.Templ
IsSystem: false,
}
if err := tmpl.Validate(); err != nil {
return model.Template{}, err
return Template{}, err
}
if err := repository.CreateTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
if err := CreateTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func listTemplates(ctx context.Context) ([]model.Template, error) {
return repository.ListTemplates(ctx)
func listTemplates(ctx context.Context) ([]Template, error) {
return ListTemplatesRecord(ctx)
}
func getTemplate(ctx context.Context, key string) (model.Template, error) {
return repository.GetTemplateByKey(ctx, key)
func getTemplate(ctx context.Context, key string) (Template, error) {
return GetTemplateByKey(ctx, key)
}
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (model.Template, error) {
tmpl, err := repository.GetTemplateByKey(ctx, key)
func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest) (Template, error) {
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return model.Template{}, err
return Template{}, err
}
tmpl.Name = req.Name
@@ -225,21 +223,21 @@ func updateTemplate(ctx context.Context, key string, req UpdateTemplateRequest)
tmpl.Content = req.Content
tmpl.Description = req.Description
if err := tmpl.Validate(); err != nil {
return model.Template{}, err
return Template{}, err
}
if err := repository.SaveTemplate(ctx, &tmpl); err != nil {
return model.Template{}, err
if err := SaveTemplateRecord(ctx, &tmpl); err != nil {
return Template{}, err
}
return tmpl, nil
}
func deleteTemplate(ctx context.Context, key string) error {
tmpl, err := repository.GetTemplateByKey(ctx, key)
tmpl, err := GetTemplateByKey(ctx, key)
if err != nil {
return err
}
if tmpl.IsSystem {
return errors.New(SystemTemplateCannotDelete)
}
return repository.DeleteTemplate(ctx, &tmpl)
return DeleteTemplateRecord(ctx, &tmpl)
}
+3 -5
View File
@@ -24,11 +24,9 @@ import (
"github.com/gin-gonic/gin"
"golang.org/x/mod/semver"
"github.com/Rain-kl/Wavelet/internal/buildinfo"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/pkg/buildinfo"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/util"
)
@@ -282,7 +280,7 @@ func (m *updaterManager) fetchRelease(ctx context.Context, repository string) (g
}
func loadRepository(ctx context.Context) (string, error) {
config, err := repository.GetSystemConfigByKey(ctx, model.ConfigKeyUpdateUpstreamRepository)
config, err := GetSystemConfigByKey(ctx, ConfigKeyUpdateUpstreamRepository)
if err != nil {
return "", fmt.Errorf("%s: %w", errInvalidRepository, err)
}
+144 -79
View File
@@ -15,11 +15,12 @@ import (
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/response"
"github.com/Rain-kl/Wavelet/pkg/util"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
)
@@ -67,7 +68,10 @@ func parseUserID(c *gin.Context) (uint64, bool) {
return id, true
}
func toUserResponse(u model.User) userResponse {
func toUserResponse(u *contracts.UserDTO) userResponse {
if u == nil {
return userResponse{}
}
return userResponse{
ID: u.ID,
Username: u.Username,
@@ -133,16 +137,16 @@ func ListUsers(c *gin.Context) {
return
}
total, modelUsers, err := listUsers(c.Request.Context(), req)
total, dtos, err := listUsers(c.Request.Context(), req)
if err != nil {
logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err)
response.AbortInternal(c, "获取用户列表失败")
return
}
users := make([]userResponse, 0, len(modelUsers))
for _, modelUser := range modelUsers {
users = append(users, toUserResponse(modelUser))
users := make([]userResponse, 0, len(dtos))
for _, dto := range dtos {
users = append(users, toUserResponse(dto))
}
c.JSON(http.StatusOK, response.OK(listUsersResponse{
@@ -243,7 +247,7 @@ func DeleteUser(c *gin.Context) {
return
}
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
@@ -334,7 +338,7 @@ func UpdateUser(c *gin.Context) {
return
}
currUser, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
currUser, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if currUser == nil {
response.AbortUnauthorized(c, AdminRequired)
return
@@ -358,40 +362,70 @@ func UpdateUser(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
func listUsers(ctx context.Context, req listUsersRequest) (int64, []model.User, error) {
return repository.ListAdminUsers(ctx, repository.AdminUserListFilter{
UserID: req.UserID,
Username: strings.TrimSpace(req.Username),
Email: strings.TrimSpace(req.Email),
Page: req.Page,
PageSize: req.PageSize,
})
func listUsers(ctx context.Context, req listUsersRequest) (int64, []*contracts.UserDTO, error) {
query := db.DB(ctx).Table("w_users")
if req.UserID != nil {
query = query.Where("id = ?", *req.UserID)
}
if req.Username != "" {
query = query.Where("username LIKE ? ESCAPE '\\'", util.EscapeLike(req.Username)+"%")
}
if req.Email != "" {
query = query.Where("email LIKE ? ESCAPE '\\'", util.EscapeLike(req.Email)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return 0, nil, err
}
var users []*contracts.UserDTO
offset := (req.Page - 1) * req.PageSize
if err := query.
Select("id, username, nickname, email, avatar_url, is_active, is_admin, last_login_at, created_at, updated_at").
Order("id ASC").
Offset(offset).
Limit(req.PageSize).
Find(&users).Error; err != nil {
return 0, nil, err
}
return total, users, nil
}
func getUserDetail(ctx context.Context, id uint64) (model.User, error) {
return repository.GetAdminUserDetail(ctx, id)
func getUserDetail(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
var user contracts.UserDTO
if err := db.DB(ctx).Table("w_users").
Select("id, username, nickname, email, avatar_url, is_active, is_admin, bio, phone, gender, website, location, last_login_at, created_at, updated_at").
Where("id = ?", id).
First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
func updateUserStatus(ctx context.Context, id uint64, active bool) error {
flags, err := repository.GetUserAdminFlags(ctx, id)
if err != nil {
var flags struct {
ID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", id).First(&flags).Error; err != nil {
return err
}
if !active && flags.IsAdmin {
return errors.New(cannotDisable)
}
var tokens []model.AccessToken
var tokenHashes []string
if !active {
tokens, _ = repository.ListAccessTokensByUserID(ctx, id)
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", id).Pluck("token_hash", &tokenHashes).Error
}
err = repository.UpdateUserActive(ctx, id, active)
err := db.DB(ctx).Table("w_users").Where("id = ?", id).Update("is_active", active).Error
if err == nil {
auth.InvalidateCachedUser(ctx, id)
if !active {
for _, token := range tokens {
auth.InvalidateCachedToken(ctx, token.TokenHash)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
}
@@ -402,77 +436,106 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
if currentUserID == targetID {
return errors.New(cannotDeleteSelf)
}
flags, err := repository.GetUserAdminFlags(ctx, targetID)
if err != nil {
var flags struct {
ID uint64
IsAdmin bool
}
if err := db.DB(ctx).Table("w_users").Select("id, is_admin").Where("id = ?", targetID).First(&flags).Error; err != nil {
return err
}
if flags.IsAdmin {
return errors.New(cannotDelete)
}
tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID)
var tokenHashes []string
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", targetID).Pluck("token_hash", &tokenHashes).Error
err = repository.DeleteUserWithRelations(ctx, targetID)
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Table("w_access_tokens").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}
if err := tx.Table("w_external_accounts").Where("user_id = ?", targetID).Delete(map[string]any{}).Error; err != nil {
return err
}
return tx.Table("w_users").Where("id = ?", targetID).Delete(map[string]any{}).Error
})
if err == nil {
auth.InvalidateCachedUser(ctx, targetID)
for _, token := range tokens {
auth.InvalidateCachedToken(ctx, token.TokenHash)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
return err
}
func createUser(ctx context.Context, req createUserRequest) (model.User, error) {
func createUser(ctx context.Context, req createUserRequest) (*contracts.UserDTO, error) {
req.Username = strings.TrimSpace(req.Username)
req.Nickname = strings.TrimSpace(req.Nickname)
req.Password = strings.TrimSpace(req.Password)
req.Email = strings.TrimSpace(req.Email)
if req.Username == "" {
return model.User{}, errors.New(usernameRequired)
return nil, errors.New(usernameRequired)
}
if req.Email == "" {
return model.User{}, errors.New(emailRequired)
return nil, errors.New(emailRequired)
}
if len(req.Password) < minPasswordLength {
return model.User{}, errors.New(passwordTooShort)
return nil, errors.New(passwordTooShort)
}
count, err := repository.CountUsersByUsername(ctx, req.Username)
if err != nil {
return model.User{}, err
var count int64
if err := db.DB(ctx).Table("w_users").Where("username = ?", req.Username).Count(&count).Error; err != nil {
return nil, err
}
if count > 0 {
return model.User{}, errors.New(usernameExists)
return nil, errors.New(usernameExists)
}
emailCount, err := repository.CountUsersByEmail(ctx, req.Email)
if err != nil {
return model.User{}, err
var emailCount int64
if err := db.DB(ctx).Table("w_users").Where("email = ?", req.Email).Count(&emailCount).Error; err != nil {
return nil, err
}
if emailCount > 0 {
return model.User{}, errors.New(emailExists)
return nil, errors.New(emailExists)
}
newUser := model.User{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
LastLoginAt: time.Time{},
hash, err := util.HashPassword(req.Password)
if err != nil {
return nil, err
}
if newUser.Nickname == "" {
newUser.Nickname = req.Username
if req.Nickname == "" {
req.Nickname = req.Username
}
if err := newUser.SetEncryptedPassword(req.Password); err != nil {
return model.User{}, err
now := time.Now()
newUser := contracts.UserDTO{
ID: idgen.NextUint64ID(),
Username: req.Username,
Nickname: req.Nickname,
Email: req.Email,
IsActive: req.IsActive,
IsAdmin: req.IsAdmin,
CreatedAt: now,
UpdatedAt: now,
}
if err := repository.CreateUser(ctx, &newUser); err != nil {
return model.User{}, err
row := map[string]any{
"id": newUser.ID,
"username": newUser.Username,
"password": hash,
"nickname": newUser.Nickname,
"email": newUser.Email,
"is_active": newUser.IsActive,
"is_admin": newUser.IsAdmin,
"created_at": now,
"updated_at": now,
}
return newUser, nil
if err := db.DB(ctx).Table("w_users").Create(row).Error; err != nil {
return nil, err
}
return &newUser, nil
}
type updateUserParam struct {
@@ -492,20 +555,18 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam
return errors.New(emailRequired)
}
targetUser, err := repository.GetAdminUserDetail(ctx, param.ID)
if err != nil {
var targetUser contracts.UserDTO
if err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).First(&targetUser).Error; err != nil {
return err
}
// 不能撤销当前登录用户的管理员权限
if currentUserID == param.ID && !param.IsAdmin && targetUser.IsAdmin {
return errors.New(cannotRevokeSelfAdmin)
}
// 如果修改了邮箱,检查邮箱是否被其他用户占用
if targetUser.Email != param.Email {
count, err := repository.CountUsersByEmail(ctx, param.Email)
if err != nil {
var count int64
if err := db.DB(ctx).Table("w_users").Where("email = ? AND id != ?", param.Email, param.ID).Count(&count).Error; err != nil {
return err
}
if count > 0 {
@@ -513,36 +574,40 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam
}
}
// 密码强度校验(如果输入了新密码)
if param.Password != "" && len(param.Password) < minPasswordLength {
return errors.New(passwordTooShort)
}
needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin)
var tokens []model.AccessToken
var tokenHashes []string
if needRevokeTokens {
tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID)
_ = db.DB(ctx).Table("w_access_tokens").Where("user_id = ?", param.ID).Pluck("token_hash", &tokenHashes).Error
}
targetUser.Nickname = param.Nickname
if targetUser.Nickname == "" {
targetUser.Nickname = targetUser.Username
if param.Nickname == "" {
param.Nickname = targetUser.Username
}
targetUser.Email = param.Email
targetUser.IsAdmin = param.IsAdmin
updates := map[string]any{
"nickname": param.Nickname,
"email": param.Email,
"is_admin": param.IsAdmin,
"updated_at": time.Now(),
}
if param.Password != "" {
if err := targetUser.SetEncryptedPassword(param.Password); err != nil {
hash, err := util.HashPassword(param.Password)
if err != nil {
return err
}
updates["password"] = hash
}
err = repository.UpdateUser(ctx, &targetUser)
err := db.DB(ctx).Table("w_users").Where("id = ?", param.ID).Updates(updates).Error
if err == nil {
auth.InvalidateCachedUser(ctx, param.ID)
if needRevokeTokens {
for _, token := range tokens {
auth.InvalidateCachedToken(ctx, token.TokenHash)
for _, hash := range tokenHashes {
auth.InvalidateCachedToken(ctx, hash)
}
}
}
+3 -3
View File
@@ -4,9 +4,9 @@
package admin
import (
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/Rain-kl/Wavelet/core/contracts"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/response"
otel_trace "github.com/Rain-kl/Wavelet/pkg/trace"
"github.com/Rain-kl/Wavelet/plugins/domain/auth"
"github.com/gin-gonic/gin"
@@ -18,7 +18,7 @@ func LoginAdminRequired() gin.HandlerFunc {
ctx, span := otel_trace.Start(c.Request.Context(), "LoginAdminRequired")
defer span.End()
user, _ := auth.GetFromContext[*model.User](c, auth.UserObjKey)
user, _ := auth.GetFromContext[*contracts.UserDTO](c, auth.UserObjKey)
if user == nil {
response.AbortNotFound(c, AdminRequired)
return
+20 -1
View File
@@ -92,6 +92,9 @@ func (Template) TableName() string {
return "w_templates"
}
// TemplateTypeEmail 邮件模板类型
const TemplateTypeEmail = "email"
// Normalize 规范化模板字段
func (t *Template) Normalize() {
t.Key = strings.TrimSpace(t.Key)
@@ -101,7 +104,7 @@ func (t *Template) Normalize() {
t.Content = strings.TrimSpace(t.Content)
t.Description = strings.TrimSpace(t.Description)
if t.Type == "" {
t.Type = "email"
t.Type = TemplateTypeEmail
}
}
@@ -201,3 +204,19 @@ type TaskExecution struct {
func (TaskExecution) TableName() string {
return "w_task_executions"
}
// ListTaskExecutionsRequest 分页查询任务执行记录请求参数
type ListTaskExecutionsRequest struct {
Page int `form:"page"`
PageSize int `form:"page_size"`
Status string `form:"status"`
TaskType string `form:"task_type"`
TaskTypes string `form:"task_types"`
TaskTypePrefix string `form:"task_type_prefix"`
}
// TaskExecutionCleanupStats 任务日志清理结果统计
type TaskExecutionCleanupStats struct {
HighFrequencyDeleted int64 `json:"high_frequency_deleted"`
LowFrequencyDeleted int64 `json:"low_frequency_deleted"`
}
+704
View File
@@ -0,0 +1,704 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/redis/go-redis/v9"
"github.com/shopspring/decimal"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/persistence/idgen"
"github.com/Rain-kl/Wavelet/pkg/util"
)
const (
configTypeSystem = "system"
errDatabaseNotInitialized = "database not initialized"
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
taskExecutionLogRedisKeyPrefix = "task:execution:log:"
taskExecutionLogExpiration = 24 * time.Hour
taskExecutionLogMaxLines = 1000
)
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
database := db.DB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
return sc, nil
}
// GetSystemConfigByGroup queries a configuration by Type and Key.
func GetSystemConfigByGroup(ctx context.Context, configType string, key string) (SystemConfig, error) {
ensureSystemConfigCacheListener()
if item, ok := ram.Get(configType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
return sc, nil
}
}
database := db.DB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
var sc SystemConfig
if err := database.Where("key = ?", key).First(&sc).Error; err != nil {
return SystemConfig{}, err
}
valBytes, err := json.Marshal(sc)
if err == nil {
ram.Set(ram.CacheItem{
Key: sc.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(sc.Key),
})
}
return sc, nil
}
// GetSystemConfigByKey queries config by key.
func GetSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
return GetSystemConfigByGroup(ctx, ConfigCacheType, key)
}
// ListSystemConfigsByKeys loads multiple config keys.
func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]SystemConfig, error) {
if len(keys) == 0 {
return map[string]SystemConfig{}, nil
}
ensureSystemConfigCacheListener()
result := make(map[string]SystemConfig, len(keys))
missing := make([]string, 0, len(keys))
for _, key := range keys {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
result[key] = sc
continue
}
}
missing = append(missing, key)
}
if len(missing) == 0 {
return result, nil
}
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Where("key IN ?", missing).Find(&configs).Error; err != nil {
return nil, err
}
for i := range configs {
valBytes, err := json.Marshal(configs[i])
if err == nil {
ram.Set(ram.CacheItem{
Key: configs[i].Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(configs[i].Key),
})
}
result[configs[i].Key] = configs[i]
}
return result, nil
}
// InvalidateVisibleSystemConfigsCache clears the cached public config list.
func InvalidateVisibleSystemConfigsCache(ctx context.Context) error {
return InvalidateAllSystemConfigCaches(ctx)
}
// ListVisibleSystemConfigs queries visible configs using local cache store.
func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
ensureSystemConfigCacheListener()
items := ram.GetTypeItems(ConfigCacheType)
if len(items) > 0 {
var list []SystemConfig
for _, item := range items {
var sc SystemConfig
if err := json.Unmarshal([]byte(item.Value), &sc); err == nil {
if sc.Visibility == ConfigVisibilityVisible {
list = append(list, sc)
}
}
}
return list, nil
}
database := db.DB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
var configs []SystemConfig
if err := database.Where("visibility = ?", ConfigVisibilityVisible).Find(&configs).Error; err != nil {
return nil, err
}
for _, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(cfg.Key),
})
}
}
return configs, nil
}
// GetIntByKey queries config and converts to int.
func GetIntByKey(ctx context.Context, key string) (int, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return 0, err
}
value, err := strconv.Atoi(sc.Value)
if err != nil {
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetDecimalByKey queries config and converts to decimal.Decimal.
func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.Decimal, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return decimal.Zero, err
}
value, err := decimal.NewFromString(sc.Value)
if err != nil {
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
}
return value.Truncate(precision), nil
}
// GetBoolByKey queries config and converts to bool.
func GetBoolByKey(ctx context.Context, key string) (bool, error) {
sc, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return false, err
}
value, err := strconv.ParseBool(sc.Value)
if err != nil {
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
}
return value, nil
}
// GetMenuDisplayConfig queries and parses menu config.
func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
sc, err := GetSystemConfigByKey(ctx, ConfigKeyMenuDisplayConfig)
if err != nil {
return nil, err
}
config := make(map[string]bool)
if sc.Value == "" || sc.Value == "{}" {
return config, nil
}
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
}
return config, nil
}
// ListAdminSystemConfigs returns all configs, optionally filtered by type.
func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemConfig, error) {
query := db.DB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
var configs []SystemConfig
if err := query.Find(&configs).Error; err != nil {
return nil, err
}
return configs, nil
}
// GetAdminSystemConfigByKey loads a config directly from DB.
func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
var config SystemConfig
if err := db.DB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return SystemConfig{}, err
}
return config, nil
}
// SystemConfigExists reports whether a config key already exists.
func SystemConfigExists(ctx context.Context, key string) (bool, error) {
var existing SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
return db.DB(ctx).Create(config).Error
}
// UpdateSystemConfigFields applies partial updates to a system config row.
func UpdateSystemConfigFields(ctx context.Context, config *SystemConfig, updates map[string]any) error {
return db.DB(ctx).Model(config).Updates(updates).Error
}
// SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache.
func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
var sc SystemConfig
err := db.DB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
if errors.Is(err, gorm.ErrRecordNotFound) {
sc = SystemConfig{
Key: key,
Value: value,
Type: configTypeSystem,
Visibility: ConfigVisibilityHidden,
}
if err := db.DB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := db.DB(ctx).Save(&sc).Error; err != nil {
return err
}
}
return InvalidateSystemConfigCache(ctx, key)
}
// ListTemplatesRecord returns all templates ordered by system flag and creation time.
func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
var templates []Template
if err := db.DB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
return nil, err
}
return templates, nil
}
// GetTemplateByKey loads a template by its key.
func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
var tmpl Template
if err := db.DB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return Template{}, err
}
return tmpl, nil
}
// TemplateExistsByKey reports whether a template key is already taken.
func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
var existing Template
err := db.DB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
if err != nil {
return false, err
}
return true, nil
}
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Delete(tmpl).Error
}
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&Schedule{}, id).Error
}
// GetScheduleByID 根据 ID 获取定时任务
func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
var schedule Schedule
if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
}
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// ListActiveSchedules 获取所有启用的定时任务
func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
}
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
}
// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetTaskExecutionByID 根据 ID 获取执行记录
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
var execution TaskExecution
if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
return nil, err
}
return &execution, nil
}
// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type.
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
var execution TaskExecution
err := db.DB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
if err == nil {
if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil {
return nil, false, loadErr
}
return &execution, true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, nil
}
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
pipe.RPush(ctx, key, line)
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
pipe.Expire(ctx, key, taskExecutionLogExpiration)
return nil
})
if err != nil {
return fmt.Errorf("append task execution log to redis: %w", err)
}
return nil
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if db.Redis == nil {
return errors.New("redis client is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
return fmt.Errorf("persist task execution log: %w", result.Error)
}
if result.RowsAffected == 0 {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := db.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
return nil
}
// ListTaskExecutionRecords 分页查询任务执行记录
func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) {
if req.Page <= 0 {
req.Page = 1
}
if req.PageSize <= 0 {
req.PageSize = 20
}
query := db.DB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
}
if req.TaskType != "" {
query = query.Where("task_type = ?", req.TaskType)
} else if types := parseTaskTypesFilter(req.TaskTypes); len(types) > 0 {
query = query.Where("task_type IN ?", types)
} else if req.TaskTypePrefix != "" {
query = query.Where("task_type LIKE ? ESCAPE '\\'", util.EscapeLike(req.TaskTypePrefix)+"%")
}
var total int64
if err := query.Count(&total).Error; err != nil {
return nil, 0, err
}
var executions []TaskExecution
offset := (req.Page - 1) * req.PageSize
if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil {
return nil, 0, err
}
if err := loadTaskExecutionLogs(ctx, executions); err != nil {
return nil, 0, err
}
return executions, total, nil
}
func parseTaskTypesFilter(raw string) []string {
if strings.TrimSpace(raw) == "" {
return nil
}
parts := strings.Split(raw, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction.
func MarkFailedTaskExecutionsSucceededTx(
tx *gorm.DB,
taskType string,
result string,
finishedAt time.Time,
) error {
return tx.Model(&TaskExecution{}).
Where("task_type = ? AND status = ?", taskType, TaskExecutionStatusFailed).
Updates(map[string]any{
"status": TaskExecutionStatusSucceeded,
"result": result,
"finished_at": finishedAt,
}).Error
}
// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention.
func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecutionCleanupStats, error) {
const (
frequencyWindowDays = 30
highFrequencyThreshold = frequencyWindowDays
)
frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays)
highFrequencyCutoff := now.AddDate(0, 0, -3)
lowFrequencyCutoff := now.AddDate(0, 0, -30)
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
Group("task_type").
Having("COUNT(*) > ?", highFrequencyThreshold).
Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err)
}
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
Delete(&TaskExecution{})
if highFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error)
}
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes)
}
lowFrequencyResult := lowFrequencyQuery.Delete(&TaskExecution{})
if lowFrequencyResult.Error != nil {
return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error)
}
return TaskExecutionCleanupStats{
HighFrequencyDeleted: highFrequencyDeleted,
LowFrequencyDeleted: lowFrequencyResult.RowsAffected,
}, nil
}
func taskExecutionLogRedisKey(taskID string) string {
return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
if db.Redis == nil {
return nil
}
logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
if db.Redis == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error {
for i := range executions {
commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1)
}
return nil
})
if err != nil {
return fmt.Errorf("get task execution logs from redis: %w", err)
}
for i := range executions {
logLines := commands[i].Val()
if len(logLines) > 0 {
executions[i].Log = strings.Join(logLines, "")
}
}
return nil
}
+208
View File
@@ -0,0 +1,208 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"encoding/json"
"errors"
"sync"
"time"
"gorm.io/gorm"
"github.com/Rain-kl/Wavelet/pkg/cache/ram"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
"github.com/Rain-kl/Wavelet/pkg/util"
)
const (
// SystemConfigBroadcastChannel broadcasts system config cache updates across nodes.
SystemConfigBroadcastChannel = "system:config_broadcast"
// SystemConfigInvalidationChannel is kept as an alias for backward compatibility.
SystemConfigInvalidationChannel = SystemConfigBroadcastChannel
// SystemConfigRedisHashKey is kept for backward compatibility in tests.
SystemConfigRedisHashKey = "system:system_configs"
// SystemConfigVisibleListRedisKey is kept for backward compatibility in tests.
SystemConfigVisibleListRedisKey = "system:visible_configs"
// ConfigCacheType is the cache type for all system configs.
ConfigCacheType = "config"
)
type systemConfigBroadcastMessage struct {
Type string `json:"type"`
Key string `json:"key"`
}
// ConfigLoader loads configuration data from the database.
type ConfigLoader struct{}
// LoadAll loads all system configs from database as CacheItems.
func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.CacheItem, error) {
configs, err := PreheatSystemConfigs(ctx)
if err != nil {
return nil, err
}
items := make([]ram.CacheItem, len(configs))
for i, cfg := range configs {
valBytes, err := json.Marshal(cfg)
if err != nil {
return nil, err
}
items[i] = ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}
}
return items, nil
}
// LoadOne loads a single system config from database as a CacheItem.
func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string) (ram.CacheItem, error) {
cfg, err := PreheatSystemConfigByKey(ctx, key)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return ram.CacheItem{}, ram.ErrNotFound
}
return ram.CacheItem{}, err
}
valBytes, err := json.Marshal(cfg)
if err != nil {
return ram.CacheItem{}, err
}
return ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: configType,
TTL: determineTTL(cfg.Key),
}, nil
}
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
func PreloadSystemConfigs(ctx context.Context) error {
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
}
var (
systemConfigListenerOnce sync.Once
systemConfigListenerCtx context.Context
systemConfigListenerCancel context.CancelFunc
systemConfigListenerDone chan struct{}
)
func ensureSystemConfigCacheListener() {
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
}
func startSystemConfigCacheInvalidationListener() {
if db.Redis == nil {
return
}
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
systemConfigListenerDone = make(chan struct{})
redisClient := db.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 db.Redis 竞争
util.Go(func() {
listenerCtx := systemConfigListenerCtx
defer close(systemConfigListenerDone)
pubsub := redisClient.Subscribe(listenerCtx, SystemConfigBroadcastChannel)
defer func() {
_ = pubsub.Close()
}()
util.Go(func() {
<-listenerCtx.Done()
_ = pubsub.Close()
})
for msg := range pubsub.Channel() {
var payload systemConfigBroadcastMessage
if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
ram.UpdateTypeItems(ConfigCacheType, nil)
continue
}
key := payload.Key
if key == "*" || key == "" {
ram.UpdateTypeItems(payload.Type, nil)
} else {
ram.Delete(payload.Type, key)
}
}
})
}
// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard.
func StopSystemConfigCacheListener() {
if systemConfigListenerCancel != nil {
systemConfigListenerCancel()
if systemConfigListenerDone != nil {
<-systemConfigListenerDone
}
systemConfigListenerCancel = nil
systemConfigListenerDone = nil
}
systemConfigListenerOnce = sync.Once{}
}
func determineTTL(_ string) time.Duration {
// Program-determined TTL: -1 means never expire for all configs by default
return -1
}
// InvalidateSystemConfigCache triggers a broadcast to refresh the cache for key.
func InvalidateSystemConfigCache(ctx context.Context, key string) error {
ensureSystemConfigCacheListener()
// Invalidate local cache synchronously first
ram.Delete(ConfigCacheType, key)
// Broadcast to other nodes and clean legacy Redis cache key
if db.Redis != nil {
_ = db.HDel(ctx, SystemConfigRedisHashKey, key)
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
}
return nil
}
// InvalidateAllSystemConfigCaches triggers a broadcast to refresh the entire config cache.
func InvalidateAllSystemConfigCaches(ctx context.Context) error {
ensureSystemConfigCacheListener()
// Invalidate all items of type ConfigCacheType synchronously first
ram.UpdateTypeItems(ConfigCacheType, nil)
// Broadcast to other nodes and clean legacy Redis cache keys
if db.Redis != nil {
_ = db.Redis.Del(ctx, db.PrefixedKey(SystemConfigRedisHashKey), db.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
}
return nil
}
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
if db.Redis == nil {
return
}
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
if err != nil {
return
}
_ = db.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
}
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
ram.ResetForTest()
}
+157
View File
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/glebarez/sqlite"
"github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications"
"gorm.io/gorm"
db "github.com/Rain-kl/Wavelet/pkg/persistence"
)
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("gorm.Open(sqlite) error = %v", err)
}
if err := sqliteDB.AutoMigrate(&SystemConfig{}); err != nil {
t.Fatalf("AutoMigrate(SystemConfig) error = %v", err)
}
siteConfig := SystemConfig{
Key: ConfigKeySiteName,
Value: "Wavelet",
Type: "system",
Description: "系统平台的展示名称",
}
if err := sqliteDB.Create(&siteConfig).Error; err != nil {
t.Fatalf("Create(site_name) error = %v", err)
}
mr, err := miniredis.Run()
if err != nil {
t.Fatalf("miniredis.Run() error = %v", err)
}
redisClient := redis.NewClient(&redis.Options{
Addr: mr.Addr(),
MaintNotificationsConfig: &maintnotifications.Config{
Mode: maintnotifications.ModeDisabled,
},
})
previousRedis := db.Redis
db.SetDB(sqliteDB)
db.Redis = redisClient
cleanup := func() {
StopSystemConfigCacheListener()
ResetSystemConfigRAMCacheForTest()
db.SetDB(nil)
db.Redis = previousRedis
_ = redisClient.Close()
mr.Close()
}
return sqliteDB, cleanup
}
func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) {
result, err := ListSystemConfigsByKeys(context.Background(), nil)
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err)
}
if len(result) != 0 {
t.Fatalf("ListSystemConfigsByKeys(nil) = %#v, want empty map", result)
}
}
func TestListSystemConfigsByKeys_LoadsFromRAMCache(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
// Initial load
warm, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err)
}
if warm.Value != "Wavelet" {
t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet")
}
// Update DB directly
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
Update("value", "db_only_value").Error; err != nil {
t.Fatalf("Update(site_name) error = %v", err)
}
// Fetch via ListSystemConfigsByKeys should serve from local store (meaning the old value "Wavelet")
configs, err := ListSystemConfigsByKeys(ctx, []string{ConfigKeySiteName})
if err != nil {
t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err)
}
sc, ok := configs[ConfigKeySiteName]
if !ok {
t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry")
}
if sc.Value != "Wavelet" {
t.Fatalf("ListSystemConfigsByKeys(site_name).Value = %q, want cached value %q", sc.Value, "Wavelet")
}
}
func TestGetSystemConfigByGroupAndInvalidation(t *testing.T) {
dbConn, cleanup := setupSystemConfigTest(t)
defer cleanup()
ctx := context.Background()
ResetSystemConfigRAMCacheForTest()
// Get via specific group/type
cfg, err := GetSystemConfigByGroup(ctx, ConfigCacheType, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByGroup error = %v", err)
}
if cfg.Value != "Wavelet" {
t.Fatalf("value = %q, want %q", cfg.Value, "Wavelet")
}
// Direct DB update
if err := dbConn.Model(&SystemConfig{}).
Where("key = ?", ConfigKeySiteName).
Update("value", "new_site_name").Error; err != nil {
t.Fatalf("DB Update error = %v", err)
}
// Invalidate
if err := InvalidateSystemConfigCache(ctx, ConfigKeySiteName); err != nil {
t.Fatalf("InvalidateSystemConfigCache error = %v", err)
}
// Wait for broadcast execution
time.Sleep(100 * time.Millisecond)
// Fetch again
updated, err := GetSystemConfigByKey(ctx, ConfigKeySiteName)
if err != nil {
t.Fatalf("GetSystemConfigByKey error = %v", err)
}
if updated.Value != "new_site_name" {
t.Fatalf("value = %q, want %q", updated.Value, "new_site_name")
}
}