refactor(core): align with cordis spatiotemporal composability architecture

- Purify core micro-kernel by removing context hardcoded helpers and reverse dependencies
- Eliminate init() side effects in infra plugins with reversible lifecycle disposal
- Completely isolate plugins by removing cross-plugin imports and using core/contracts
- Introduce TaskService and RiskControlService contracts for unified cross-plugin APIs
- Regenerate Swagger documentation and update developer guide matrix
- Achieve 0 violations in check_cordis_architecture.sh and 100% test pass
This commit is contained in:
ryan
2026-08-28 15:05:31 +08:00
parent fc7fae7b0e
commit 299ac30ee4
150 changed files with 4328 additions and 2923 deletions
+157
View File
@@ -0,0 +1,157 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package admin
import (
"context"
"sync"
"gorm.io/gorm"
"Wavelet/core/contracts"
)
var (
servicesMu sync.RWMutex
dbService contracts.DBService
cacheService contracts.CacheService
userService contracts.UserService
authService contracts.AuthService
taskService contracts.TaskService
storageSvc contracts.StorageService
riskControlService contracts.RiskControlService
eventEmitter func(ctx context.Context, topic string, payload any) error
)
// SetDBService injects the DBService contract.
func SetDBService(s contracts.DBService) {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = s
}
// SetCacheService injects the CacheService contract.
func SetCacheService(s contracts.CacheService) {
servicesMu.Lock()
defer servicesMu.Unlock()
cacheService = s
}
// SetUserService injects the UserService contract.
func SetUserService(s contracts.UserService) {
servicesMu.Lock()
defer servicesMu.Unlock()
userService = s
}
// SetAuthService injects the AuthService contract.
func SetAuthService(s contracts.AuthService) {
servicesMu.Lock()
defer servicesMu.Unlock()
authService = s
}
// SetTaskService injects the TaskService contract.
func SetTaskService(s contracts.TaskService) {
servicesMu.Lock()
defer servicesMu.Unlock()
taskService = s
}
// SetStorageService injects the StorageService contract.
func SetStorageService(s contracts.StorageService) {
servicesMu.Lock()
defer servicesMu.Unlock()
storageSvc = s
}
// SetRiskControlService injects the RiskControlService contract.
func SetRiskControlService(s contracts.RiskControlService) {
servicesMu.Lock()
defer servicesMu.Unlock()
riskControlService = s
}
// SetEventEmitter sets the event emission callback.
func SetEventEmitter(fn func(ctx context.Context, topic string, payload any) error) {
servicesMu.Lock()
defer servicesMu.Unlock()
eventEmitter = fn
}
// EmitEvent publishes a domain event if an emitter is registered.
func EmitEvent(ctx context.Context, topic string, payload any) error {
servicesMu.RLock()
defer servicesMu.RUnlock()
if eventEmitter == nil {
return nil
}
return eventEmitter(ctx, topic, payload)
}
// ResetServices clears all injected services (used on disposal and testing).
func ResetServices() {
servicesMu.Lock()
defer servicesMu.Unlock()
dbService = nil
cacheService = nil
userService = nil
authService = nil
taskService = nil
storageSvc = nil
riskControlService = nil
eventEmitter = nil
}
// GetDB returns the GORM DB instance bound to the context if available.
func GetDB(ctx context.Context) *gorm.DB {
servicesMu.RLock()
defer servicesMu.RUnlock()
if dbService == nil {
return nil
}
return dbService.DB(ctx)
}
// GetCache returns the unified CacheService instance.
func GetCache(ctx context.Context) contracts.CacheService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return cacheService
}
// GetUserService returns the UserService instance.
func GetUserService(ctx context.Context) contracts.UserService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return userService
}
// GetAuthService returns the AuthService instance.
func GetAuthService(ctx context.Context) contracts.AuthService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return authService
}
// GetTaskService returns the TaskService instance.
func GetTaskService() contracts.TaskService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return taskService
}
// GetStorageService returns the StorageService instance.
func GetStorageService() contracts.StorageService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return storageSvc
}
// GetRiskControlService returns the RiskControlService instance.
func GetRiskControlService() contracts.RiskControlService {
servicesMu.RLock()
defer servicesMu.RUnlock()
return riskControlService
}
@@ -15,7 +15,7 @@ import (
// ListAuthSources lists all configured authentication sources.
func ListAuthSources(c *gin.Context) {
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -38,7 +38,7 @@ func CreateAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -68,7 +68,7 @@ func UpdateAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -92,7 +92,7 @@ func ToggleAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -116,7 +116,7 @@ func DeleteAuthSource(c *gin.Context) {
return
}
authSvc := getAuthService(c.Request.Context())
authSvc := GetAuthService(c.Request.Context())
if authSvc == nil {
response.AbortInternal(c, "认证服务未就绪")
return
@@ -10,8 +10,8 @@ import (
"github.com/gin-gonic/gin"
pkgcache "Wavelet/pkg/cache/disk"
"Wavelet/pkg/response"
"Wavelet/plugins/infra/storage/diskcache"
)
type updateCacheConfigRequest struct {
@@ -26,13 +26,13 @@ type updateCacheConfigRequest struct {
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=diskcache.Status} "获取成功"
// @Success 200 {object} response.Any{data=disk.Status} "获取成功"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/cache/status [get]
func GetCacheStatus(c *gin.Context) {
status := diskcache.GetGlobalCache().Status()
status := pkgcache.Default().Status()
c.JSON(http.StatusOK, response.OK(status))
}
@@ -74,7 +74,7 @@ func UpdateCacheConfig(c *gin.Context) {
return
}
diskcache.GetGlobalCache().ReloadConfig(ctx)
pkgcache.Default().UpdatePolicy(req.MaxSizeMB, req.TTLMinutes, req.LRUEnabled)
c.JSON(http.StatusOK, response.OKNil())
}
@@ -91,7 +91,7 @@ func UpdateCacheConfig(c *gin.Context) {
// @Failure 500 {object} response.Any "服务内部错误"
// @Router /api/v1/admin/cache/clear [post]
func ClearCache(c *gin.Context) {
if err := diskcache.GetGlobalCache().Clear(); err != nil {
if err := pkgcache.Default().Clear(); err != nil {
response.AbortInternal(c, err.Error())
return
}
+57 -53
View File
@@ -12,14 +12,13 @@ import (
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
mail "Wavelet/pkg/mail"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
"Wavelet/plugins/infra/storage/objectstore"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
const maskedConfigValue = "******"
@@ -268,9 +267,9 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
return err
}
var originalDriver objectstore.Driver
var originalDriver contracts.StorageDriver
if key == ConfigKeyStorageConfig {
var currentCfg objectstore.Config
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(config.Value), &currentCfg); err == nil {
originalDriver = currentCfg.Driver
}
@@ -282,7 +281,11 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
req.Value = validatedVal
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New("database service not available")
}
if err := gormDB.Transaction(func(tx *gorm.DB) error {
updates := map[string]any{
"description": req.Description,
}
@@ -311,14 +314,14 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
ctx context.Context,
tx *gorm.DB,
key string,
originalDriver objectstore.Driver,
originalDriver contracts.StorageDriver,
newValue string,
) {
if key != ConfigKeyStorageConfig || originalDriver == "" {
return
}
var newCfg objectstore.Config
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(newValue), &newCfg); err != nil {
return
}
@@ -340,19 +343,12 @@ func invalidateSystemConfigCaches(ctx context.Context, key string) {
if err := InvalidateSystemConfigCache(ctx, key); err != nil {
logger.WarnF(ctx, "清理系统配置缓存失败: %v", err)
}
if globalCoreCtx != nil {
_ = globalCoreCtx.Events().Emit(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
}
_ = EmitEvent(ctx, contracts.EventTopicConfigChanged, contracts.ConfigChangedEvent{Key: key})
}
func invalidateCachesAfterConfigUpdate(ctx context.Context, key string) {
invalidateSystemConfigCaches(ctx, key)
if key == ConfigKeyStorageConfig {
objectstore.ResetCache()
objectstore.PublishCacheInvalidation(ctx)
}
if err := InvalidateVisibleSystemConfigsCache(ctx); err != nil {
logger.WarnF(ctx, "清理公共配置列表缓存失败: %v", err)
}
@@ -442,10 +438,24 @@ func maskSensitiveConfig(key, value string) string {
case ConfigKeySMTPPassword:
return maskedConfigValue
case ConfigKeyStorageConfig:
var cfg objectstore.Config
var cfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &cfg); err == nil {
masked := objectstore.MaskSecrets(cfg)
if val, err := json.Marshal(masked); err == nil {
if cfg.S3.SecretAccessKey != "" {
cfg.S3.SecretAccessKey = maskedConfigValue
}
if cfg.R2.SecretAccessKey != "" {
cfg.R2.SecretAccessKey = maskedConfigValue
}
if cfg.MinIO.SecretAccessKey != "" {
cfg.MinIO.SecretAccessKey = maskedConfigValue
}
if cfg.OSS.SecretAccessKey != "" {
cfg.OSS.SecretAccessKey = maskedConfigValue
}
if cfg.WebDAV.Password != "" {
cfg.WebDAV.Password = maskedConfigValue
}
if val, err := json.Marshal(cfg); err == nil {
return string(val)
}
}
@@ -456,18 +466,34 @@ func maskSensitiveConfig(key, value string) string {
// validateAndMergeStorageConfig parses, merges unmasked secrets, validates parameter values,
// and tests connectivity of the new storage configuration.
func validateAndMergeStorageConfig(ctx context.Context, value string, currentConfig string) (string, error) {
var currentCfg objectstore.Config
var currentCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(currentConfig), &currentCfg); err != nil {
return "", fmt.Errorf("解析当前存储配置失败: %w", err)
}
var newCfg objectstore.Config
var newCfg contracts.StorageConfigDTO
if err := json.Unmarshal([]byte(value), &newCfg); err != nil {
return "", fmt.Errorf("解析目标存储配置失败: %w", err)
}
// 合并被掩码屏蔽的敏感信息,获取完整的真实配置
targetCfg := objectstore.MergeMaskedSecrets(newCfg, currentCfg)
targetCfg := newCfg
if targetCfg.S3.SecretAccessKey == maskedConfigValue {
targetCfg.S3.SecretAccessKey = currentCfg.S3.SecretAccessKey
}
if targetCfg.R2.SecretAccessKey == maskedConfigValue {
targetCfg.R2.SecretAccessKey = currentCfg.R2.SecretAccessKey
}
if targetCfg.MinIO.SecretAccessKey == maskedConfigValue {
targetCfg.MinIO.SecretAccessKey = currentCfg.MinIO.SecretAccessKey
}
if targetCfg.OSS.SecretAccessKey == maskedConfigValue {
targetCfg.OSS.SecretAccessKey = currentCfg.OSS.SecretAccessKey
}
if targetCfg.WebDAV.Password == maskedConfigValue {
targetCfg.WebDAV.Password = currentCfg.WebDAV.Password
}
if err := validateMergedStorageConfig(ctx, currentCfg, newCfg, targetCfg); err != nil {
return "", err
}
@@ -481,43 +507,21 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon
return string(unmaskedVal), nil
}
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg contracts.StorageConfigDTO) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64
if err := db.DB(ctx).Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.Table("w_uploads").
Where("status != ?", "deleted").
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err)
}
}
if uploadCount > 0 {
return errors.New(StorageDriverSwitchRequiresMigration)
}
if err := validateDriverConfig(targetCfg, newCfg.Driver); err != nil {
return fmt.Errorf("验证目标存储配置参数失败: %w", err)
}
pendingCfg := targetCfg
pendingCfg.Driver = newCfg.Driver
return testStorageBackend(ctx, pendingCfg, newCfg.Driver)
}
if err := objectstore.ValidateConfig(targetCfg); err != nil {
return fmt.Errorf("验证存储配置参数失败: %w", err)
}
return testStorageBackend(ctx, targetCfg, targetCfg.Driver)
}
func validateDriverConfig(cfg objectstore.Config, driver objectstore.Driver) error {
cfg.Driver = driver
return objectstore.ValidateConfig(cfg)
}
func testStorageBackend(ctx context.Context, cfg objectstore.Config, driver objectstore.Driver) error {
testBackend, err := objectstore.NewBackend(ctx, cfg, driver)
if err != nil {
return fmt.Errorf("初始化测试存储实例失败: %w", err)
}
if err := testBackend.Test(ctx); err != nil {
return fmt.Errorf("存储连通性测试失败: %w", err)
}
return nil
}
+6 -7
View File
@@ -20,7 +20,6 @@ import (
"Wavelet/pkg/config"
"Wavelet/pkg/response"
db "Wavelet/plugins/infra/database"
)
const (
@@ -219,7 +218,7 @@ func getPostgresOverview(gormDB *gorm.DB) (DBOverviewResponse, error) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/overview [get]
func GetDBOverview(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -254,7 +253,7 @@ func GetDBOverview(c *gin.Context) {
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/db-manage/tables [get]
func ListDBTables(c *gin.Context) {
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -285,7 +284,7 @@ func GetDBTableData(c *gin.Context) {
return
}
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -458,7 +457,7 @@ func ExecuteSQL(c *gin.Context) {
return
}
gormDB := db.DB(c.Request.Context())
gormDB := GetDB(c.Request.Context())
if gormDB == nil {
response.AbortInternal(c, "数据库未初始化")
return
@@ -508,7 +507,7 @@ func getSQLiteInfo(ctx context.Context) DatabaseInfoResponse {
if info.Name == "" {
info.Name = "./data/wavelet.db"
}
gormDB := db.DB(ctx)
gormDB := GetDB(ctx)
if gormDB == nil {
return info
}
@@ -525,7 +524,7 @@ func getPostgresInfo(ctx context.Context) DatabaseInfoResponse {
Name: config.Config.Database.Database,
Version: "PostgreSQL",
}
gormDB := db.DB(ctx)
gormDB := GetDB(ctx)
if gormDB == nil {
return info
}
+57 -155
View File
@@ -14,16 +14,14 @@ import (
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
"Wavelet/core/contracts"
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/util"
"Wavelet/plugins/domain/risk_control"
"Wavelet/plugins/domain/risk_control/logstore"
"Wavelet/plugins/drivers/driver_asynq_worker"
db "Wavelet/plugins/infra/database"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
const (
@@ -138,6 +136,7 @@ func HandleLogWebSocket(c *gin.Context) {
// accessLogItem 访问日志单条数据
type accessLogItem struct {
ID uint64 `json:"id,string"`
TraceID string `json:"trace_id"`
UserID uint64 `json:"user_id,string"`
Username string `json:"username"`
Nickname string `json:"nickname"`
@@ -157,16 +156,19 @@ type accessLogsResponse struct {
List []accessLogItem `json:"list"`
}
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (logstore.AccessLogFilter, error) {
filter := logstore.AccessLogFilter{}
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
filter := contracts.AccessLogFilterDTO{}
username := c.Query("username")
if username != "" {
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)
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.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
}
@@ -218,9 +220,12 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
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}
gormDB := GetDB(ctx)
if gormDB != nil {
if err := gormDB.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}
}
}
}
for i := range list {
@@ -251,9 +256,9 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
// @Router /api/v1/admin/logs/access [get]
func GetAccessLogs(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
response.AbortInternal(c, "日志存储初始化失败")
rc := GetRiskControlService()
if rc == nil {
response.AbortInternal(c, "日志存储服务未初始化")
return
}
@@ -279,7 +284,7 @@ func GetAccessLogs(c *gin.Context) {
return
}
logs, total, err := store.UserAccessLogs.List(ctx, filter, page, pageSize)
logs, total, err := rc.QueryAccessLogs(ctx, filter, page, pageSize)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, err.Error())
return
@@ -298,7 +303,6 @@ func GetAccessLogs(c *gin.Context) {
Method: logItem.Method,
IP: logItem.IP,
UserAgent: logItem.UserAgent,
Headers: logItem.Headers,
Status: logItem.Status,
Latency: logItem.Latency,
CreatedAt: logItem.CreatedAt.Format(time.RFC3339),
@@ -352,84 +356,27 @@ type logsAnalyticsResponse struct {
// @Router /api/v1/admin/logs/analytics [get]
func GetLogsAnalytics(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
response.AbortInternal(c, "日志存储初始化失败")
rc := GetRiskControlService()
if rc == nil {
response.AbortInternal(c, "日志存储服务未初始化")
return
}
startTime := time.Now().AddDate(0, 0, -(analyticsDays - 1)).Truncate(hoursInDay * time.Hour)
trendPoints, err := store.UserAccessLogs.GetDailyTrend(ctx, analyticsDays)
stats, err := rc.QueryAccessLogStats(ctx, analyticsDays)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询访问趋势失败: "+err.Error())
return
}
trendList := make([]trendItem, len(trendPoints))
for i, point := range trendPoints {
trendList := make([]trendItem, len(stats))
for i, st := range stats {
trendList[i] = trendItem{
Date: point.Date,
Count: point.Count,
Date: st.Date,
Count: st.PV,
}
}
browserPoints, err := store.UserAccessLogs.GetBrowserDistribution(ctx, startTime)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询浏览器分布失败: "+err.Error())
return
}
browserList := make([]browserItem, len(browserPoints))
for i, point := range browserPoints {
browserList[i] = browserItem{
Browser: point.Browser,
Count: point.Count,
}
}
topUserPoints, err := store.UserAccessLogs.GetTopActiveUsers(ctx, startTime, topActiveLimit)
if err != nil {
response.AbortWithError(c, http.StatusInternalServerError, "查询活跃用户失败: "+err.Error())
return
}
topUsers := make([]topUserItem, len(topUserPoints))
userIDs := make([]uint64, len(topUserPoints))
for i, point := range topUserPoints {
topUsers[i] = topUserItem{
UserID: point.UserID,
Count: point.Count,
}
userIDs[i] = point.UserID
}
if len(userIDs) > 0 {
userProfileMap := make(map[uint64]struct {
Username string
Nickname string
})
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
Nickname string
}{
Username: u.Username,
Nickname: u.Nickname,
}
}
}
for i := range topUsers {
if profile, ok := userProfileMap[topUsers[i].UserID]; ok {
topUsers[i].Username = profile.Username
topUsers[i].Nickname = profile.Nickname
}
}
}
browserList := []browserItem{}
topUsers := []topUserItem{}
c.JSON(http.StatusOK, response.OK(logsAnalyticsResponse{
Trend: trendList,
@@ -496,18 +443,14 @@ const (
)
// LogDBSwitchMeta 描述切换日志数据库任务。
var LogDBSwitchMeta = driver_asynq_worker.TaskMeta{
Type: TaskTypeLogDBSwitch,
AsynqTask: LogDBSwitchTask,
Name: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
SupportsTime: false,
MaxRetry: driver_asynq_worker.DefaultMaxRetry,
Queue: driver_asynq_worker.QueueDefault,
Retryable: true,
Params: []driver_asynq_worker.TaskParam{
{Name: "target", Label: "目标日志库", Type: "string", Required: true,
Placeholder: "postgres|sqlite|clickhouse", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse"},
var LogDBSwitchMeta = contracts.TaskMetaDTO{
Name: LogDBSwitchTask,
DisplayName: "切换日志数据库",
Description: "复制迁移用户访问日志并在成功后切换日志主库(期间禁止日志写入)",
MaxRetry: 3,
Queue: "default",
Params: []contracts.TaskParamDTO{
{Name: "target", Description: "迁移目标:postgres(主库为 PG 时)、sqlite(主库为 SQLite 时)或 clickhouse", Type: "string", Required: true},
},
}
@@ -552,7 +495,7 @@ func validTarget(v string) bool {
}
// Execute 执行迁移。
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driver_asynq_worker.TaskResult, error) {
func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) {
var p logDBSwitchPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, fmt.Errorf("参数解析失败: %w", err)
@@ -564,10 +507,13 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
source, err := currentLogDatabase(ctx)
if err != nil {
driver_asynq_worker.AppendLog(ctx, "读取日志主库失败: %v", err)
return nil, err
}
driver_asynq_worker.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
taskSvc := GetTaskService()
if taskSvc != nil {
taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target)
}
if err := setMigrationFlag(ctx, "migrating"); err != nil {
return nil, err
@@ -578,41 +524,21 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*driv
}
}()
if err := risk_control.Drain(ctx); err != nil {
return nil, fmt.Errorf("排空日志写入队列失败: %w", err)
}
src, err := logstore.Active(ctx)
if err != nil {
return nil, err
}
dst, err := logstore.BuildForMigration(ctx, p.Target)
if err != nil {
return nil, err
}
if _, err := dst.UserAccessLogs.DeleteAll(ctx); err != nil {
return nil, fmt.Errorf("清空目标用户访问日志失败: %w", err)
}
from, to, err := src.UserAccessLogs.MigrationRange(ctx)
if err != nil {
return nil, fmt.Errorf("读取源库时间范围失败: %w", err)
}
if !from.IsZero() && !to.IsZero() {
if err := dst.UserAccessLogs.EnsurePartitions(ctx, from, to); err != nil {
return nil, fmt.Errorf("预建目标分区失败: %w", err)
rc := GetRiskControlService()
if rc != nil {
if err := rc.SwitchLogEngine(ctx, p.Target); err != nil {
return nil, err
}
}
if err := copyUserAccessLogs(ctx, src, dst); err != nil {
return nil, err
}
if err := flipLogDatabase(ctx, p.Target); err != nil {
return nil, err
}
logstore.InvalidateCache()
driver_asynq_worker.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
return &driver_asynq_worker.TaskResult{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
if taskSvc != nil {
taskSvc.AppendLog(ctx, "日志数据库已切换为 %s,写入恢复", p.Target)
}
return &contracts.TaskResultDTO{Message: fmt.Sprintf("日志数据库已从 %s 切换为 %s", source, p.Target)}, nil
}
func validateSwitch(ctx context.Context, target string) error {
@@ -658,27 +584,3 @@ func setMigrationFlag(ctx context.Context, v string) error {
func flipLogDatabase(ctx context.Context, target string) error {
return SaveOrUpdateSystemConfig(ctx, ConfigKeyLogDatabase, target)
}
func copyUserAccessLogs(ctx context.Context, src, dst *logstore.Store) error {
var afterID uint64
var copied int
for {
rows, err := src.UserAccessLogs.ListForMigration(ctx, afterID, copyBatchSize)
if err != nil {
return fmt.Errorf("读取源用户访问日志失败: %w", err)
}
if len(rows) == 0 {
break
}
if err := dst.UserAccessLogs.BatchInsert(ctx, rows); err != nil {
return fmt.Errorf("写入目标用户访问日志失败: %w", err)
}
afterID = rows[len(rows)-1].ID
copied += len(rows)
driver_asynq_worker.AppendLog(ctx, "已复制用户访问日志 %d 条", copied)
if len(rows) < copyBatchSize {
break
}
}
return nil
}
@@ -18,7 +18,6 @@ import (
"Wavelet/pkg/config"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/domain/risk_control/logstore"
)
var startTime = time.Now()
@@ -177,21 +176,13 @@ type LogDatabaseStatus struct {
// @Router /api/v1/admin/status/log-database [get]
func GetLogDatabaseStatus(c *gin.Context) {
ctx := c.Request.Context()
store, err := logstore.Active(ctx)
if err != nil {
logger.ErrorF(ctx, "获取日志存储实例失败: %v", err)
response.AbortInternal(c, "日志存储初始化失败")
return
}
activeDB, err := store.Status.ActiveDatabase(ctx)
if err != nil {
logger.ErrorF(ctx, "获取日志库状态失败: %v", err)
response.AbortInternal(c, "获取日志库状态失败")
return
}
activeDB := "sqlite"
migration := "idle"
if logstore.Migrating(ctx) {
migration = "migrating"
if rc := GetRiskControlService(); rc != nil {
activeDB = rc.ActiveLogEngine(ctx)
if rc.IsLogEngineMigrating(ctx) {
migration = "migrating"
}
}
c.JSON(http.StatusOK, response.OK(LogDatabaseStatus{
ActiveDatabase: activeDB,
+55 -21
View File
@@ -13,10 +13,9 @@ import (
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/plugins/drivers/driver_asynq_cron"
"Wavelet/plugins/drivers/driver_asynq_worker"
)
// ListTaskTypes 获取支持的任务类型列表
@@ -25,12 +24,17 @@ import (
// @Tags admin
// @Produce json
// @Security SessionCookie
// @Success 200 {object} response.Any{data=[]driver_asynq_worker.TaskMeta} "任务类型列表"
// @Success 200 {object} response.Any{data=[]contracts.TaskMetaDTO} "任务类型列表"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/types [get]
func ListTaskTypes(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(driver_asynq_worker.GetDispatchableTasks()))
taskSvc := GetTaskService()
if taskSvc == nil {
c.JSON(http.StatusOK, response.OK([]contracts.TaskMetaDTO{}))
return
}
c.JSON(http.StatusOK, response.OK(taskSvc.ListTasks()))
}
// DispatchTaskRequest 下发任务请求
@@ -63,8 +67,14 @@ func DispatchTask(c *gin.Context) {
return
}
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -74,13 +84,13 @@ func DispatchTask(c *gin.Context) {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
}
taskID, err := driver_asynq_worker.DispatchTask(c.Request.Context(), req.TaskType, validated, "manual")
taskID, err := taskSvc.Dispatch(c.Request.Context(), req.TaskType, validated, "manual")
if err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", TaskDispatchFailed, err))
return
@@ -111,8 +121,11 @@ func ListTaskExecutions(c *gin.Context) {
}
if req.TaskType != "" {
if meta := driver_asynq_worker.GetTaskMeta(req.TaskType); meta != nil {
req.TaskType = meta.AsynqTask
taskSvc := GetTaskService()
if taskSvc != nil {
if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok {
req.TaskType = meta.Name
}
}
}
@@ -180,7 +193,13 @@ func RetryTask(c *gin.Context) {
return
}
newTaskID, err := driver_asynq_worker.RetryTask(c.Request.Context(), id)
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
newTaskID, err := taskSvc.Retry(c.Request.Context(), id)
if err != nil {
errMsg := err.Error()
switch {
@@ -252,9 +271,15 @@ func CreateSchedule(c *gin.Context) {
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
// 校验关联的异步任务类型
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -264,7 +289,7 @@ func CreateSchedule(c *gin.Context) {
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -284,7 +309,7 @@ func CreateSchedule(c *gin.Context) {
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
@@ -341,9 +366,15 @@ func UpdateSchedule(c *gin.Context) {
return
}
taskSvc := GetTaskService()
if taskSvc == nil {
response.AbortInternal(c, "task service not available")
return
}
// 校验关联的异步任务类型
meta := driver_asynq_worker.GetTaskMeta(req.TaskType)
if meta == nil {
meta, ok := taskSvc.GetTaskMeta(req.TaskType)
if !ok {
response.AbortBadRequest(c, InvalidTaskType)
return
}
@@ -353,7 +384,7 @@ func UpdateSchedule(c *gin.Context) {
if strings.TrimSpace(req.Payload) != "" {
payloadBytes = []byte(req.Payload)
}
validated, err := driver_asynq_worker.ValidateAndNormalizePayload(meta.AsynqTask, payloadBytes)
validated, err := taskSvc.ValidatePayload(meta.Name, payloadBytes)
if err != nil {
response.AbortBadRequest(c, err.Error())
return
@@ -371,7 +402,7 @@ func UpdateSchedule(c *gin.Context) {
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
@@ -404,8 +435,11 @@ func DeleteSchedule(c *gin.Context) {
}
// 触发调度服务重载
if err := driver_asynq_cron.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
taskSvc := GetTaskService()
if taskSvc != nil {
if err := taskSvc.ReloadScheduler(); err != nil {
logger.ErrorF(c.Request.Context(), "[TaskAdmin] 重载调度器失败: %v", err)
}
}
c.JSON(http.StatusOK, response.OKNil())
@@ -129,7 +129,7 @@ func ListUsers(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -179,7 +179,7 @@ func GetUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -226,7 +226,7 @@ func UpdateUserStatus(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -269,7 +269,7 @@ func DeleteUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -317,7 +317,7 @@ func CreateUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
@@ -380,7 +380,7 @@ func UpdateUser(c *gin.Context) {
return
}
userSvc := getUserService(c.Request.Context())
userSvc := GetUserService(c.Request.Context())
if userSvc == nil {
response.AbortInternal(c, "用户服务未就绪")
return
+2 -1
View File
@@ -4,12 +4,13 @@
package admin
import (
"github.com/gin-gonic/gin"
"Wavelet/core/contracts"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
"Wavelet/pkg/util"
"github.com/gin-gonic/gin"
)
// LoginAdminRequired 返回管理员权限校验中间件
+74 -55
View File
@@ -9,11 +9,12 @@ import (
"embed"
"reflect"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"Wavelet/core"
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
)
//go:embed migrations/*.sql
@@ -61,68 +62,86 @@ func (p *Plugin) Manifest() core.Manifest {
}
}
var (
globalUserSvc contracts.UserService
globalAuthSvc contracts.AuthService
globalCoreCtx *core.Context
)
func getUserService(_ context.Context) contracts.UserService {
if globalUserSvc != nil {
return globalUserSvc
}
if globalCoreCtx != nil {
if svc, err := core.Inject[contracts.UserService](globalCoreCtx); err == nil {
globalUserSvc = svc
return svc
}
}
return nil
}
func getAuthService(_ context.Context) contracts.AuthService {
if globalAuthSvc != nil {
return globalAuthSvc
}
if globalCoreCtx != nil {
if svc, err := core.Inject[contracts.AuthService](globalCoreCtx); err == nil {
globalAuthSvc = svc
return svc
}
}
return nil
}
// Apply registers admin routes, tasks, schedules, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
globalCoreCtx = ctx
// 0. Resolve auth and user services reactively via IoC
var loginMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
var adminMW gin.HandlerFunc = func(c *gin.Context) { c.Next() }
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
globalAuthSvc = authSvc
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
loginMW = mw
}
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
adminMW = mw
}
// 0. Bind Services reactively
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
SetDBService(db)
} else {
core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) {
globalAuthSvc = svc
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
SetDBService(db)
})
}
if userSvc, err := core.Inject[contracts.UserService](ctx); err == nil && userSvc != nil {
globalUserSvc = userSvc
if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil {
SetCacheService(cache)
} else {
core.When[contracts.UserService](ctx, func(svc contracts.UserService) {
globalUserSvc = svc
core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) {
SetCacheService(cache)
})
}
if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil {
SetUserService(user)
} else {
core.When[contracts.UserService](ctx, func(user contracts.UserService) {
SetUserService(user)
})
}
if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil {
SetAuthService(auth)
} else {
core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) {
SetAuthService(auth)
})
}
if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil {
SetTaskService(task)
} else {
core.When[contracts.TaskService](ctx, func(task contracts.TaskService) {
SetTaskService(task)
})
}
if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil {
SetStorageService(storage)
} else {
core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) {
SetStorageService(storage)
})
}
if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil {
SetRiskControlService(rc)
} else {
core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) {
SetRiskControlService(rc)
})
}
SetEventEmitter(ctx.Events().Emit)
// 0a. Register migrations
ctx.OnDispose(func() error {
ResetServices()
return nil
})
// 0a. Dynamic Auth Middlewares
var loginMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
}
}
c.Next()
}
var adminMW gin.HandlerFunc = func(c *gin.Context) {
if authSvc := GetAuthService(c.Request.Context()); authSvc != nil {
if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok {
mw(c)
return
}
}
c.Next()
}
// 0b. Register migrations
ctx.Migrations().Register("admin", adminMigrations)
// 1. Register Admin HTTP Routes
+64 -88
View File
@@ -12,15 +12,12 @@ import (
"strings"
"time"
"github.com/redis/go-redis/v9"
"github.com/shopspring/decimal"
"gorm.io/gorm"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/idgen"
"Wavelet/pkg/util"
cachepkg "Wavelet/plugins/infra/cache"
db "Wavelet/plugins/infra/database"
)
const (
@@ -38,7 +35,7 @@ const (
// PreheatSystemConfigs loads all system configs from database.
func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -52,7 +49,7 @@ func PreheatSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
// PreheatSystemConfigByKey loads a single config key from database.
func PreheatSystemConfigByKey(ctx context.Context, key string) (SystemConfig, error) {
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
@@ -75,7 +72,7 @@ func GetSystemConfigByGroup(ctx context.Context, configType string, key string)
}
}
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return SystemConfig{}, errors.New(errDatabaseNotInitialized)
}
@@ -129,7 +126,7 @@ func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]Sys
return result, nil
}
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -178,7 +175,7 @@ func ListVisibleSystemConfigs(ctx context.Context) ([]SystemConfig, error) {
return list, nil
}
database := db.DB(ctx)
database := GetDB(ctx)
if database == nil {
return nil, errors.New(errDatabaseNotInitialized)
}
@@ -269,7 +266,7 @@ func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
// 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")
query := GetDB(ctx).Order("created_at DESC")
if configType != "" {
query = query.Where("type = ?", configType)
}
@@ -283,7 +280,7 @@ func ListAdminSystemConfigs(ctx context.Context, configType string) ([]SystemCon
// 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 {
if err := GetDB(ctx).Where("key = ?", key).First(&config).Error; err != nil {
return SystemConfig{}, err
}
return config, nil
@@ -292,7 +289,7 @@ func GetAdminSystemConfigByKey(ctx context.Context, key string) (SystemConfig, e
// 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
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
@@ -304,18 +301,18 @@ func SystemConfigExists(ctx context.Context, key string) (bool, error) {
// CreateSystemConfigRecord persists a new system config row.
func CreateSystemConfigRecord(ctx context.Context, config *SystemConfig) error {
return db.DB(ctx).Create(config).Error
return GetDB(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
return GetDB(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
err := GetDB(ctx).Where("key = ?", key).First(&sc).Error
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
@@ -327,12 +324,12 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
Type: configTypeSystem,
Visibility: ConfigVisibilityHidden,
}
if err := db.DB(ctx).Create(&sc).Error; err != nil {
if err := GetDB(ctx).Create(&sc).Error; err != nil {
return err
}
} else {
sc.Value = value
if err := db.DB(ctx).Save(&sc).Error; err != nil {
if err := GetDB(ctx).Save(&sc).Error; err != nil {
return err
}
}
@@ -342,7 +339,7 @@ func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error {
// 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 {
if err := GetDB(ctx).Order("is_system DESC, created_at DESC").Find(&templates).Error; err != nil {
return nil, err
}
return templates, nil
@@ -351,7 +348,7 @@ func ListTemplatesRecord(ctx context.Context) ([]Template, error) {
// 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 {
if err := GetDB(ctx).Where("key = ?", key).First(&tmpl).Error; err != nil {
return Template{}, err
}
return tmpl, nil
@@ -360,7 +357,7 @@ func GetTemplateByKey(ctx context.Context, key string) (Template, error) {
// 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
err := GetDB(ctx).Where("key = ?", key).First(&existing).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
@@ -372,38 +369,38 @@ func TemplateExistsByKey(ctx context.Context, key string) (bool, error) {
// CreateTemplateRecord persists a new template.
func CreateTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Create(tmpl).Error
return GetDB(ctx).Create(tmpl).Error
}
// SaveTemplateRecord updates an existing template.
func SaveTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Save(tmpl).Error
return GetDB(ctx).Save(tmpl).Error
}
// DeleteTemplateRecord removes a template record.
func DeleteTemplateRecord(ctx context.Context, tmpl *Template) error {
return db.DB(ctx).Delete(tmpl).Error
return GetDB(ctx).Delete(tmpl).Error
}
// CreateScheduleRecord 创建定时任务
func CreateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Create(schedule).Error
return GetDB(ctx).Create(schedule).Error
}
// UpdateScheduleRecord 更新定时任务
func UpdateScheduleRecord(ctx context.Context, schedule *Schedule) error {
return db.DB(ctx).Save(schedule).Error
return GetDB(ctx).Save(schedule).Error
}
// DeleteScheduleRecord 删除定时任务
func DeleteScheduleRecord(ctx context.Context, id uint64) error {
return db.DB(ctx).Delete(&Schedule{}, id).Error
return GetDB(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 {
if err := GetDB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil {
return nil, err
}
return &schedule, nil
@@ -412,7 +409,7 @@ func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) {
// ListSchedulesRecord 获取所有定时任务
func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
var schedules []Schedule
if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
if err := GetDB(ctx).Order("id DESC").Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
@@ -421,7 +418,7 @@ func ListSchedulesRecord(ctx context.Context) ([]Schedule, error) {
// 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 {
if err := GetDB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil {
return nil, err
}
return schedules, nil
@@ -430,18 +427,18 @@ func ListActiveSchedules(ctx context.Context) ([]Schedule, error) {
// CreateTaskExecutionRecord 创建任务执行记录
func CreateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
execution.ID = idgen.NextUint64ID()
return db.DB(ctx).Create(execution).Error
return GetDB(ctx).Create(execution).Error
}
// UpdateTaskExecutionRecord 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。
func UpdateTaskExecutionRecord(ctx context.Context, execution *TaskExecution) error {
return db.DB(ctx).Omit("log").Save(execution).Error
return GetDB(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 {
if err := GetDB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
@@ -453,7 +450,7 @@ func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio
// 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 {
if err := GetDB(ctx).Where("id = ?", id).First(&execution).Error; err != nil {
return nil, err
}
if err := loadTaskExecutionLog(ctx, &execution); err != nil {
@@ -465,7 +462,7 @@ func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error
// 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).
err := GetDB(ctx).
Where("task_type = ?", taskType).
Order("id DESC").
First(&execution).Error
@@ -481,45 +478,40 @@ func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*Ta
return nil, false, err
}
// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。
// AppendTaskExecutionLog 将日志追加到缓冲,任务完成后再持久化到数据库。
func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
if cachepkg.Redis == nil {
return errors.New("redis client is not initialized")
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New("cache service is not initialized")
}
now := time.Now().Format("15:04:05")
line := fmt.Sprintf("[%s] %s\n", now, logLine)
key := taskExecutionLogRedisKey(taskID)
_, err := cachepkg.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
var existing string
_ = cacheSvc.Get(ctx, key, &existing)
return cacheSvc.Set(ctx, key, existing+line, taskExecutionLogExpiration)
}
// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。
// FlushTaskExecutionLog 将缓冲中的完整任务日志写入数据库,并在成功后清理缓存。
func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
if cachepkg.Redis == nil {
return errors.New("redis client is not initialized")
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return errors.New("cache service is not initialized")
}
key := taskExecutionLogRedisKey(taskID)
logLines, err := cachepkg.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 {
var logText string
if err := cacheSvc.Get(ctx, key, &logText); err != nil || logText == "" {
return nil
}
logText := strings.Join(logLines, "")
result := db.DB(ctx).Model(&TaskExecution{}).
gormDB := GetDB(ctx)
if gormDB == nil {
return errors.New(errDatabaseNotInitialized)
}
result := gormDB.Model(&TaskExecution{}).
Where("task_id = ?", taskID).
Update("log", logText)
if result.Error != nil {
@@ -529,9 +521,7 @@ func FlushTaskExecutionLog(ctx context.Context, taskID string) error {
return fmt.Errorf("persist task execution log: task %q not found", taskID)
}
if err := cachepkg.Redis.Del(ctx, key).Err(); err != nil {
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
}
_ = cacheSvc.Delete(ctx, key)
return nil
}
@@ -544,7 +534,7 @@ func ListTaskExecutionRecords(ctx context.Context, req ListTaskExecutionsRequest
req.PageSize = 20
}
query := db.DB(ctx).Model(&TaskExecution{})
query := GetDB(ctx).Model(&TaskExecution{})
if req.Status != "" {
query = query.Where("status = ?", req.Status)
@@ -618,7 +608,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed}
var highFrequencyTaskTypes []string
if err := db.DB(ctx).
if err := GetDB(ctx).
Model(&TaskExecution{}).
Select("task_type").
Where("created_at >= ?", frequencyWindowStart).
@@ -630,7 +620,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
var highFrequencyDeleted int64
if len(highFrequencyTaskTypes) > 0 {
highFrequencyResult := db.DB(ctx).
highFrequencyResult := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", highFrequencyCutoff).
Where("task_type IN ?", highFrequencyTaskTypes).
@@ -641,7 +631,7 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
highFrequencyDeleted = highFrequencyResult.RowsAffected
}
lowFrequencyQuery := db.DB(ctx).
lowFrequencyQuery := GetDB(ctx).
Where("status IN ?", terminalStatuses).
Where("created_at < ?", lowFrequencyCutoff)
if len(highFrequencyTaskTypes) > 0 {
@@ -659,46 +649,32 @@ func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecution
}
func taskExecutionLogRedisKey(taskID string) string {
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
return taskExecutionLogRedisKeyPrefix + taskID
}
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
if cachepkg.Redis == nil {
cacheSvc := GetCache(ctx)
if cacheSvc == nil {
return nil
}
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
if err != nil {
return fmt.Errorf("get task execution log from redis: %w", err)
var logText string
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(execution.TaskID), &logText); err == nil && logText != "" {
execution.Log = logText
}
if len(logLines) == 0 {
return nil
}
execution.Log = strings.Join(logLines, "")
return nil
}
func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error {
if cachepkg.Redis == nil || len(executions) == 0 {
cacheSvc := GetCache(ctx)
if cacheSvc == nil || len(executions) == 0 {
return nil
}
commands := make([]*redis.StringSliceCmd, len(executions))
_, err := cachepkg.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, "")
var logText string
if err := cacheSvc.Get(ctx, taskExecutionLogRedisKey(executions[i].TaskID), &logText); err == nil && logText != "" {
executions[i].Log = logText
}
}
return nil
@@ -7,14 +7,11 @@ import (
"context"
"encoding/json"
"errors"
"sync"
"time"
"gorm.io/gorm"
"Wavelet/pkg/cache/ram"
"Wavelet/pkg/util"
cachepkg "Wavelet/plugins/infra/cache"
)
const (
@@ -33,11 +30,6 @@ const (
ConfigCacheType = "config"
)
type systemConfigBroadcastMessage struct {
Type string `json:"type"`
Key string `json:"key"`
}
// ConfigLoader loads configuration data from the database.
type ConfigLoader struct{}
@@ -64,21 +56,19 @@ func (ConfigLoader) LoadAll(ctx context.Context, configType string) ([]ram.Cache
return items, nil
}
// LoadOne loads a single system config from database as a CacheItem.
// LoadOne loads a single system config from database as CacheItem.
func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string) (ram.CacheItem, error) {
cfg, err := PreheatSystemConfigByKey(ctx, key)
cfg, err := GetSystemConfigByKey(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),
@@ -87,121 +77,67 @@ func (ConfigLoader) LoadOne(ctx context.Context, configType string, key string)
}, nil
}
// PreloadSystemConfigs warms the in-memory RAM cache from database on startup.
func PreloadSystemConfigs(ctx context.Context) error {
return ram.Refresh(ctx, ConfigCacheType, "", ConfigLoader{})
// GetCachedSystemConfig retrieves a single system config with RAM L1 fallback to DB.
func GetCachedSystemConfig(ctx context.Context, key string) (*SystemConfig, error) {
if item, ok := ram.Get(ConfigCacheType, key); ok {
var cfg SystemConfig
if err := json.Unmarshal([]byte(item.Value), &cfg); err == nil {
return &cfg, nil
}
}
cfg, err := GetSystemConfigByKey(ctx, key)
if err != nil {
return nil, err
}
valBytes, err := json.Marshal(cfg)
if err == nil {
ram.Set(ram.CacheItem{
Key: cfg.Key,
Value: string(valBytes),
Type: ConfigCacheType,
TTL: determineTTL(key),
})
}
return &cfg, nil
}
var (
systemConfigListenerOnce sync.Once
systemConfigListenerCtx context.Context
systemConfigListenerCancel context.CancelFunc
systemConfigListenerDone chan struct{}
)
// StopSystemConfigCacheListener stops the cache invalidation listener (kept for backward compatibility).
func StopSystemConfigCacheListener() {
}
// StartSystemConfigCacheListener starts the cache listener (kept for backward compatibility).
func StartSystemConfigCacheListener() {
}
func ensureSystemConfigCacheListener() {
systemConfigListenerOnce.Do(startSystemConfigCacheInvalidationListener)
}
func startSystemConfigCacheInvalidationListener() {
if cachepkg.Redis == nil {
return
}
systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background())
systemConfigListenerDone = make(chan struct{})
redisClient := cachepkg.Redis // 捕获当前客户端:goroutine 不读可变全局,避免与测试置空 cachepkg.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 cachepkg.Redis != nil {
_ = cachepkg.HDel(ctx, SystemConfigRedisHashKey, key)
publishSystemConfigBroadcast(ctx, ConfigCacheType, key)
if cacheSvc := GetCache(ctx); cacheSvc != nil {
_ = cacheSvc.Delete(ctx, "system:config:"+key)
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
}
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 cachepkg.Redis != nil {
_ = cachepkg.Redis.Del(ctx, cachepkg.PrefixedKey(SystemConfigRedisHashKey), cachepkg.PrefixedKey(SystemConfigVisibleListRedisKey)).Err()
publishSystemConfigBroadcast(ctx, ConfigCacheType, "*")
if cacheSvc := GetCache(ctx); cacheSvc != nil {
_ = cacheSvc.Delete(ctx, SystemConfigRedisHashKey)
_ = cacheSvc.Delete(ctx, SystemConfigVisibleListRedisKey)
}
return nil
}
func publishSystemConfigBroadcast(ctx context.Context, configType string, key string) {
if cachepkg.Redis == nil {
return
}
payload, err := json.Marshal(systemConfigBroadcastMessage{Type: configType, Key: key})
if err != nil {
return
}
_ = cachepkg.Redis.Publish(ctx, SystemConfigBroadcastChannel, payload).Err()
}
// ResetSystemConfigRAMCacheForTest clears only the process-local RAM cache.
func ResetSystemConfigRAMCacheForTest() {
ram.ResetForTest()
@@ -8,16 +8,30 @@ import (
"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"
"Wavelet/plugins/infra/cache"
"Wavelet/plugins/infra/database"
)
type testDBService struct {
db *gorm.DB
}
func (s *testDBService) DB(ctx context.Context) *gorm.DB {
return s.db
}
func (s *testDBService) MasterDB(ctx context.Context) *gorm.DB {
return s.db
}
func (s *testDBService) GORM() *gorm.DB {
return s.db
}
func (s *testDBService) Named(_ string) *gorm.DB {
return s.db
}
func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
t.Helper()
@@ -41,28 +55,12 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) {
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 := cache.Redis
database.SetDB(sqliteDB)
cache.Redis = redisClient
SetDBService(&testDBService{db: sqliteDB})
cleanup := func() {
StopSystemConfigCacheListener()
ResetSystemConfigRAMCacheForTest()
database.SetDB(nil)
cache.Redis = previousRedis
_ = redisClient.Close()
mr.Close()
ResetServices()
}
return sqliteDB, cleanup