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
@@ -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()