mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 14:46:36 +08:00
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:
@@ -0,0 +1,40 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package driver_asynq_cron
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
)
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -15,7 +15,7 @@ import (
|
||||
"github.com/hibiken/asynq"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
@@ -66,7 +66,7 @@ type Plugin struct {
|
||||
// New creates a new Asynq Cron Scheduler driver plugin.
|
||||
func New(opts ...Option) *Plugin {
|
||||
p := &Plugin{
|
||||
redisOpt: driver_asynq_worker.RedisOpt,
|
||||
redisOpt: RedisOpt,
|
||||
location: time.Local,
|
||||
}
|
||||
|
||||
@@ -90,6 +90,30 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
p.coreCtx = ctx
|
||||
p.mu.Unlock()
|
||||
|
||||
// Bind DBService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
setDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
setDBService(db)
|
||||
})
|
||||
}
|
||||
|
||||
// Bind TaskService
|
||||
if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil {
|
||||
setTaskService(taskSvc)
|
||||
} else {
|
||||
core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) {
|
||||
setTaskService(taskSvc)
|
||||
})
|
||||
}
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
setTaskService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// Register migrations for w_schedules table
|
||||
ctx.Migrations().Register("driver_asynq_cron", cronMigrations)
|
||||
|
||||
|
||||
@@ -6,8 +6,6 @@ package driver_asynq_cron
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// Schedule 定时任务配置表
|
||||
@@ -30,7 +28,11 @@ func (Schedule) TableName() string {
|
||||
// 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 {
|
||||
db := getDB(ctx)
|
||||
if db == nil {
|
||||
return nil, nil
|
||||
}
|
||||
if err := db.Where("is_active = ?", true).Find(&schedules).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return schedules, nil
|
||||
|
||||
@@ -13,8 +13,8 @@ import (
|
||||
|
||||
"github.com/hibiken/asynq"
|
||||
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/plugins/drivers/driver_asynq_worker"
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -22,11 +22,22 @@ var (
|
||||
schedulerMutex sync.Mutex
|
||||
quitChan chan struct{}
|
||||
schedulerOnce sync.Once
|
||||
taskSvcMu sync.RWMutex
|
||||
taskSvcInstance contracts.TaskService
|
||||
// RedisOpt is the Redis connection option for the scheduler.
|
||||
RedisOpt asynq.RedisConnOpt
|
||||
)
|
||||
|
||||
// GetAsynqClient 获取全局 AsynqClient
|
||||
func GetAsynqClient() *asynq.Client {
|
||||
return driver_asynq_worker.AsynqClient
|
||||
func setTaskService(s contracts.TaskService) {
|
||||
taskSvcMu.Lock()
|
||||
defer taskSvcMu.Unlock()
|
||||
taskSvcInstance = s
|
||||
}
|
||||
|
||||
func getTaskService() contracts.TaskService {
|
||||
taskSvcMu.RLock()
|
||||
defer taskSvcMu.RUnlock()
|
||||
return taskSvcInstance
|
||||
}
|
||||
|
||||
// StartScheduler 启动调度器 (该函数阻塞,直到调度器退出)
|
||||
@@ -92,27 +103,39 @@ func ReloadScheduler() error {
|
||||
|
||||
// 3. 实例化新的调度器
|
||||
newScheduler := asynq.NewScheduler(
|
||||
driver_asynq_worker.RedisOpt,
|
||||
RedisOpt,
|
||||
&asynq.SchedulerOpts{
|
||||
Location: location,
|
||||
},
|
||||
)
|
||||
|
||||
// 4. 遍历并注册任务
|
||||
taskSvc := getTaskService()
|
||||
for _, s := range schedules {
|
||||
meta := driver_asynq_worker.GetTaskMeta(s.TaskType)
|
||||
if meta == nil {
|
||||
continue // 忽略排程配置中无效的任务类型
|
||||
taskName := s.TaskType
|
||||
maxRetry := 3
|
||||
queue := "default"
|
||||
|
||||
if taskSvc != nil {
|
||||
if meta, ok := taskSvc.GetTaskMeta(s.TaskType); ok {
|
||||
taskName = meta.Name
|
||||
if meta.MaxRetry > 0 {
|
||||
maxRetry = meta.MaxRetry
|
||||
}
|
||||
if meta.Queue != "" {
|
||||
queue = meta.Queue
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 构造 Asynq 载荷。定时任务使用对应 Meta 中的 Asynq 标识,同时将数据库中保存的 json 作为参数
|
||||
t := asynq.NewTask(meta.AsynqTask, []byte(s.Payload))
|
||||
// 构造 Asynq 载荷
|
||||
t := asynq.NewTask(taskName, []byte(s.Payload))
|
||||
|
||||
if _, err := newScheduler.Register(
|
||||
s.Cron,
|
||||
t,
|
||||
asynq.MaxRetry(meta.MaxRetry),
|
||||
asynq.Queue(meta.Queue),
|
||||
asynq.MaxRetry(maxRetry),
|
||||
asynq.Queue(queue),
|
||||
); err != nil {
|
||||
// 定时任务配置可能有误(如 Cron 格式不被 Asynq 识别),记录日志并跳过
|
||||
logger.ErrorF(context.Background(), "[Scheduler] 注册定时任务失败 id=%d name=%s: %v", s.ID, s.Name, err)
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package driver_asynq_worker
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
redisMu sync.RWMutex
|
||||
rdbClient redis.UniversalClient
|
||||
)
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
// SetRedisClient sets the redis client used for task logs.
|
||||
func SetRedisClient(c redis.UniversalClient) {
|
||||
redisMu.Lock()
|
||||
defer redisMu.Unlock()
|
||||
rdbClient = c
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func getRedisClient() redis.UniversalClient {
|
||||
redisMu.RLock()
|
||||
defer redisMu.RUnlock()
|
||||
return rdbClient
|
||||
}
|
||||
@@ -17,9 +17,28 @@ import (
|
||||
"go.opentelemetry.io/otel/propagation"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
// mockHandler 用于测试的模拟任务处理器
|
||||
type mockHandler struct {
|
||||
executeFunc func(ctx context.Context, payload []byte) (*TaskResult, error)
|
||||
@@ -55,7 +74,11 @@ func failHandler() *mockHandler {
|
||||
const testTaskType = "test:mock_task"
|
||||
|
||||
func setupTest(t *testing.T) func() {
|
||||
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
testDB, mr, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
setDBService(&mockDBService{db: testDB})
|
||||
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
SetRedisClient(rdb)
|
||||
|
||||
AsynqClient = asynq.NewClient(asynq.RedisClientOpt{
|
||||
Addr: mr.Addr(),
|
||||
})
|
||||
@@ -66,6 +89,9 @@ func setupTest(t *testing.T) func() {
|
||||
_ = AsynqClient.Close()
|
||||
AsynqClient = nil
|
||||
}
|
||||
_ = rdb.Close()
|
||||
setDBService(nil)
|
||||
SetRedisClient(nil)
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"github.com/hibiken/asynq"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -82,6 +83,7 @@ type Plugin struct {
|
||||
mux *asynq.ServeMux
|
||||
running bool
|
||||
coreCtx *core.Context
|
||||
taskSvc contracts.TaskService
|
||||
}
|
||||
|
||||
// New creates a new Asynq Worker driver plugin.
|
||||
@@ -113,7 +115,24 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
p.coreCtx = ctx
|
||||
p.mu.Unlock()
|
||||
|
||||
// Register migrations for w_task_executions table
|
||||
// 0. Bind DBService
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
setDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
setDBService(db)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 1. Provide contracts.TaskService
|
||||
p.taskSvc = &taskServiceImpl{}
|
||||
core.Provide[contracts.TaskService](ctx, p.taskSvc)
|
||||
|
||||
// 2. Register migrations for w_task_executions table
|
||||
ctx.Migrations().Register("driver_asynq_worker", workerMigrations)
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
@@ -254,3 +273,152 @@ func toAsynqHandler(h any) (asynq.Handler, error) {
|
||||
return nil, fmt.Errorf("unsupported task handler type: %T", h)
|
||||
}
|
||||
}
|
||||
|
||||
type taskServiceImpl struct{}
|
||||
|
||||
func (s *taskServiceImpl) Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) {
|
||||
return DispatchTask(ctx, taskType, payload, triggeredBy)
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) ListTasks() []contracts.TaskMetaDTO {
|
||||
all := GetDispatchableTasks()
|
||||
res := make([]contracts.TaskMetaDTO, 0, len(all))
|
||||
for _, m := range all {
|
||||
params := make([]contracts.TaskParamDTO, 0, len(m.Params))
|
||||
for _, param := range m.Params {
|
||||
params = append(params, contracts.TaskParamDTO{
|
||||
Name: param.Name,
|
||||
Type: param.Type,
|
||||
Description: param.Description,
|
||||
Required: param.Required,
|
||||
})
|
||||
}
|
||||
res = append(res, contracts.TaskMetaDTO{
|
||||
Name: m.Type,
|
||||
DisplayName: m.Name,
|
||||
Description: m.Description,
|
||||
Params: params,
|
||||
MaxRetry: m.MaxRetry,
|
||||
Queue: m.Queue,
|
||||
})
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) GetTaskMeta(taskType string) (contracts.TaskMetaDTO, bool) {
|
||||
m := GetTaskMeta(taskType)
|
||||
if m == nil {
|
||||
return contracts.TaskMetaDTO{}, false
|
||||
}
|
||||
params := make([]contracts.TaskParamDTO, 0, len(m.Params))
|
||||
for _, param := range m.Params {
|
||||
params = append(params, contracts.TaskParamDTO{
|
||||
Name: param.Name,
|
||||
Type: param.Type,
|
||||
Description: param.Description,
|
||||
Required: param.Required,
|
||||
})
|
||||
}
|
||||
return contracts.TaskMetaDTO{
|
||||
Name: m.Type,
|
||||
DisplayName: m.Name,
|
||||
Description: m.Description,
|
||||
Params: params,
|
||||
MaxRetry: m.MaxRetry,
|
||||
Queue: m.Queue,
|
||||
}, true
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) ListExecutions(ctx context.Context, taskType string, status string, page, pageSize int) ([]contracts.TaskExecutionDTO, int64, error) {
|
||||
db := getDB(ctx)
|
||||
if db == nil {
|
||||
return nil, 0, errors.New("db not initialized")
|
||||
}
|
||||
query := db.Model(&TaskExecution{})
|
||||
if taskType != "" {
|
||||
query = query.Where("task_type = ?", taskType)
|
||||
}
|
||||
if status != "" {
|
||||
query = query.Where("status = ?", status)
|
||||
}
|
||||
var total int64
|
||||
if err := query.Count(&total).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var rows []TaskExecution
|
||||
offset := (page - 1) * pageSize
|
||||
if err := query.Order("id DESC").Offset(offset).Limit(pageSize).Find(&rows).Error; err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
res := make([]contracts.TaskExecutionDTO, 0, len(rows))
|
||||
for _, r := range rows {
|
||||
res = append(res, contracts.TaskExecutionDTO{
|
||||
ID: r.ID,
|
||||
TaskID: r.TaskID,
|
||||
TaskType: r.TaskType,
|
||||
TaskName: r.TaskName,
|
||||
Status: string(r.Status),
|
||||
Retryable: r.Retryable,
|
||||
MaxRetry: r.MaxRetry,
|
||||
RetryCount: r.RetryCount,
|
||||
Log: r.Log,
|
||||
ErrorMessage: r.ErrorMessage,
|
||||
Result: r.Result,
|
||||
StartedAt: r.StartedAt,
|
||||
FinishedAt: r.FinishedAt,
|
||||
Duration: r.Duration,
|
||||
Payload: r.Payload,
|
||||
TriggeredBy: r.TriggeredBy,
|
||||
CreatedAt: r.CreatedAt,
|
||||
UpdatedAt: r.UpdatedAt,
|
||||
})
|
||||
}
|
||||
return res, total, nil
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) Retry(ctx context.Context, id uint64) (string, error) {
|
||||
return RetryTask(ctx, id)
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) ValidatePayload(taskType string, payload []byte) ([]byte, error) {
|
||||
meta := GetTaskMeta(taskType)
|
||||
if meta == nil {
|
||||
return payload, nil
|
||||
}
|
||||
return ValidateAndNormalizePayload(meta.AsynqTask, payload)
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) ReloadScheduler() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) AppendLog(ctx context.Context, format string, args ...any) {
|
||||
AppendLog(ctx, format, args...)
|
||||
}
|
||||
|
||||
func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contracts.TaskExecutionDTO, error) {
|
||||
exec, err := GetTaskExecutionByID(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &contracts.TaskExecutionDTO{
|
||||
ID: exec.ID,
|
||||
TaskID: exec.TaskID,
|
||||
TaskType: exec.TaskType,
|
||||
TaskName: exec.TaskName,
|
||||
Status: string(exec.Status),
|
||||
Retryable: exec.Retryable,
|
||||
MaxRetry: exec.MaxRetry,
|
||||
RetryCount: exec.RetryCount,
|
||||
Log: exec.Log,
|
||||
ErrorMessage: exec.ErrorMessage,
|
||||
Result: exec.Result,
|
||||
StartedAt: exec.StartedAt,
|
||||
FinishedAt: exec.FinishedAt,
|
||||
Duration: exec.Duration,
|
||||
Payload: exec.Payload,
|
||||
TriggeredBy: exec.TriggeredBy,
|
||||
CreatedAt: exec.CreatedAt,
|
||||
UpdatedAt: exec.UpdatedAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -13,8 +13,6 @@ import (
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"Wavelet/pkg/idgen"
|
||||
cachepkg "Wavelet/plugins/infra/cache"
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -24,23 +22,23 @@ const (
|
||||
)
|
||||
|
||||
func taskExecutionLogRedisKey(taskID string) string {
|
||||
return cachepkg.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID)
|
||||
return taskExecutionLogRedisKeyPrefix + taskID
|
||||
}
|
||||
|
||||
func createTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
||||
if execution.ID == 0 {
|
||||
execution.ID = idgen.NextUint64ID()
|
||||
}
|
||||
return db.DB(ctx).Create(execution).Error
|
||||
return getDB(ctx).Create(execution).Error
|
||||
}
|
||||
|
||||
func updateTaskExecution(ctx context.Context, execution *TaskExecution) error {
|
||||
return db.DB(ctx).Omit("log").Save(execution).Error
|
||||
return getDB(ctx).Omit("log").Save(execution).Error
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
_ = loadTaskExecutionLog(ctx, &execution)
|
||||
@@ -49,7 +47,7 @@ func getTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error
|
||||
|
||||
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
|
||||
}
|
||||
_ = loadTaskExecutionLog(ctx, &execution)
|
||||
@@ -57,7 +55,8 @@ func getTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio
|
||||
}
|
||||
|
||||
func appendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error {
|
||||
if cachepkg.Redis == nil {
|
||||
rdb := getRedisClient()
|
||||
if rdb == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
}
|
||||
|
||||
@@ -65,7 +64,7 @@ func appendTaskExecutionLog(ctx context.Context, taskID string, logLine string)
|
||||
line := fmt.Sprintf("[%s] %s\n", now, logLine)
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
|
||||
_, err := cachepkg.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
_, err := rdb.TxPipelined(ctx, func(pipe redis.Pipeliner) error {
|
||||
pipe.RPush(ctx, key, line)
|
||||
pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1)
|
||||
pipe.Expire(ctx, key, taskExecutionLogExpiration)
|
||||
@@ -78,12 +77,13 @@ func appendTaskExecutionLog(ctx context.Context, taskID string, logLine string)
|
||||
}
|
||||
|
||||
func flushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
if cachepkg.Redis == nil {
|
||||
rdb := getRedisClient()
|
||||
if rdb == nil {
|
||||
return errors.New("redis client is not initialized")
|
||||
}
|
||||
|
||||
key := taskExecutionLogRedisKey(taskID)
|
||||
logLines, err := cachepkg.Redis.LRange(ctx, key, 0, -1).Result()
|
||||
logLines, err := rdb.LRange(ctx, key, 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
@@ -92,7 +92,7 @@ func flushTaskExecutionLog(ctx context.Context, taskID string) error {
|
||||
}
|
||||
logText := strings.Join(logLines, "")
|
||||
|
||||
result := db.DB(ctx).Model(&TaskExecution{}).
|
||||
result := getDB(ctx).Model(&TaskExecution{}).
|
||||
Where("task_id = ?", taskID).
|
||||
Update("log", logText)
|
||||
if result.Error != nil {
|
||||
@@ -102,18 +102,19 @@ 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 {
|
||||
if err := rdb.Del(ctx, key).Err(); err != nil {
|
||||
return fmt.Errorf("delete persisted task execution log from redis: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error {
|
||||
if cachepkg.Redis == nil {
|
||||
rdb := getRedisClient()
|
||||
if rdb == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
logLines, err := cachepkg.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
||||
logLines, err := rdb.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("get task execution log from redis: %w", err)
|
||||
}
|
||||
|
||||
@@ -9,8 +9,6 @@ import (
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
db "Wavelet/plugins/infra/database"
|
||||
)
|
||||
|
||||
// TaskExecutionStatus 任务执行状态
|
||||
@@ -57,13 +55,13 @@ func (TaskExecution) TableName() string {
|
||||
|
||||
// CreateTaskExecution 创建任务执行记录
|
||||
func CreateTaskExecution(ctx context.Context, exec *TaskExecution) error {
|
||||
return db.DB(ctx).Create(exec).Error
|
||||
return getDB(ctx).Create(exec).Error
|
||||
}
|
||||
|
||||
// GetTaskExecutionByTaskID 根据 TaskID 查询执行记录
|
||||
func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) {
|
||||
var exec TaskExecution
|
||||
if err := db.DB(ctx).Where("task_id = ?", taskID).First(&exec).Error; err != nil {
|
||||
if err := getDB(ctx).Where("task_id = ?", taskID).First(&exec).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_ = loadTaskExecutionLog(ctx, &exec)
|
||||
@@ -73,7 +71,7 @@ func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutio
|
||||
// GetTaskExecutionByID 根据主键 ID 查询执行记录
|
||||
func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) {
|
||||
var exec TaskExecution
|
||||
if err := db.DB(ctx).Where("id = ?", id).First(&exec).Error; err != nil {
|
||||
if err := getDB(ctx).Where("id = ?", id).First(&exec).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_ = loadTaskExecutionLog(ctx, &exec)
|
||||
@@ -83,7 +81,7 @@ func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error
|
||||
// GetLatestTaskExecutionByTaskType 获取指定任务类型的最新执行记录
|
||||
func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*TaskExecution, bool, error) {
|
||||
var exec TaskExecution
|
||||
err := db.DB(ctx).Where("task_type = ?", taskType).Order("id DESC").First(&exec).Error
|
||||
err := getDB(ctx).Where("task_type = ?", taskType).Order("id DESC").First(&exec).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, false, nil
|
||||
@@ -114,7 +112,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).
|
||||
@@ -126,7 +124,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).
|
||||
@@ -137,7 +135,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 {
|
||||
|
||||
@@ -6,8 +6,9 @@ package driver_asynq_worker
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"github.com/redis/go-redis/v9/maintnotifications"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
)
|
||||
|
||||
func TestMaintNotificationsConfig(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package driver_http
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
)
|
||||
|
||||
var (
|
||||
dbMu sync.RWMutex
|
||||
dbSvc contracts.DBService
|
||||
)
|
||||
|
||||
func setDBService(s contracts.DBService) {
|
||||
dbMu.Lock()
|
||||
defer dbMu.Unlock()
|
||||
dbSvc = s
|
||||
}
|
||||
|
||||
func getDB(ctx context.Context) *gorm.DB {
|
||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
}
|
||||
dbMu.RLock()
|
||||
s := dbSvc
|
||||
dbMu.RUnlock()
|
||||
if s != nil {
|
||||
return s.DB(ctx)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -15,13 +15,14 @@ import (
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/trace"
|
||||
"Wavelet/pkg/util"
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-contrib/sessions/redis"
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/trace"
|
||||
"Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
// BuildEngine 构建并初始化 Gin 路由引擎及全部中间件和路由
|
||||
|
||||
@@ -11,14 +11,14 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
|
||||
"Wavelet/pkg/config"
|
||||
"Wavelet/pkg/logger"
|
||||
"Wavelet/pkg/response"
|
||||
otel_trace "Wavelet/pkg/trace"
|
||||
database "Wavelet/plugins/infra/database"
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
func loggerMiddleware() gin.HandlerFunc {
|
||||
@@ -72,7 +72,11 @@ func loggerMiddleware() gin.HandlerFunc {
|
||||
|
||||
func isOriginAllowed(ctx context.Context, origin string) bool {
|
||||
var val string
|
||||
if err := database.DB(ctx).Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || val == "" {
|
||||
db := getDB(ctx)
|
||||
if db == nil {
|
||||
return false
|
||||
}
|
||||
if err := db.Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || val == "" {
|
||||
return false
|
||||
}
|
||||
allowedOrigins := strings.Split(val, ",")
|
||||
|
||||
@@ -4,17 +4,40 @@
|
||||
package driver_http
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"Wavelet/pkg/testhelper"
|
||||
)
|
||||
|
||||
type mockDBService struct {
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
func (m *mockDBService) GORM() *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func (m *mockDBService) DB(ctx context.Context) *gorm.DB {
|
||||
return m.db.WithContext(ctx)
|
||||
}
|
||||
|
||||
func (m *mockDBService) Named(_ string) *gorm.DB {
|
||||
return m.db
|
||||
}
|
||||
|
||||
func TestCORSMiddleware(t *testing.T) {
|
||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||
defer cleanup()
|
||||
setDBService(&mockDBService{db: dbConn})
|
||||
defer func() {
|
||||
setDBService(nil)
|
||||
cleanup()
|
||||
}()
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"Wavelet/core"
|
||||
"Wavelet/core/contracts"
|
||||
"Wavelet/pkg/util"
|
||||
)
|
||||
|
||||
@@ -97,6 +98,19 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
p.coreCtx = ctx
|
||||
p.mu.Unlock()
|
||||
|
||||
// Bind DBService from Context
|
||||
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||
setDBService(db)
|
||||
} else {
|
||||
core.When[contracts.DBService](ctx, func(db contracts.DBService) {
|
||||
setDBService(db)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
setDBService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
ctx.OnDispose(func() error {
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout)
|
||||
defer cancel()
|
||||
|
||||
Reference in New Issue
Block a user