From 3a61c38dc5e6f8f66de1eec89a15e17a518d4de7 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 28 Aug 2026 16:09:18 +0800 Subject: [PATCH] fix(core): delegate driver and task extension lookups to root context and improve test reliability --- .gitignore | 1 + backend/cmd/app.go | 10 +++-- backend/cmd/redis_plug_test.go | 9 +++- backend/core/app_test.go | 8 ++-- backend/core/context.go | 28 ++++++------ .../drivers/driver_asynq_worker/executor.go | 4 +- .../drivers/driver_asynq_worker/plugin.go | 45 +++++++++---------- .../drivers/driver_asynq_worker/utils.go | 36 +++++++++++++-- 8 files changed, 89 insertions(+), 52 deletions(-) diff --git a/.gitignore b/.gitignore index e552cd90..e34e3322 100644 --- a/.gitignore +++ b/.gitignore @@ -65,3 +65,4 @@ s3_cache /.superpowers/ /backend/plugins/domain/upload/filesrv/uploads/ /backend/plugins/domain/upload/task/uploads/ +/backend/data/ diff --git a/backend/cmd/app.go b/backend/cmd/app.go index dae69d79..4383fee3 100644 --- a/backend/cmd/app.go +++ b/backend/cmd/app.go @@ -113,12 +113,16 @@ type sharedStore struct { func (s *sharedStore) Tablename() string { return "w_schema_versions" } func (s *sharedStore) CreateVersionTable(ctx context.Context, db goosedb.DBTxConn) error { - _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS w_schema_versions ( + timeType := "TIMESTAMPTZ" + if s.dialect == "sqlite3" || s.dialect == "sqlite" { + timeType = "DATETIME" + } + _, err := db.ExecContext(ctx, fmt.Sprintf(`CREATE TABLE IF NOT EXISTS w_schema_versions ( plugin_id VARCHAR(64) NOT NULL, version_id BIGINT NOT NULL, - applied_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + applied_at %s NOT NULL DEFAULT CURRENT_TIMESTAMP, PRIMARY KEY (plugin_id, version_id) - )`) + )`, timeType)) return err } diff --git a/backend/cmd/redis_plug_test.go b/backend/cmd/redis_plug_test.go index 184fa433..d4be231d 100644 --- a/backend/cmd/redis_plug_test.go +++ b/backend/cmd/redis_plug_test.go @@ -21,7 +21,12 @@ import ( func TestRedisPluggability_Simulation(t *testing.T) { origRedisEnabled := config.Config.Redis.Enabled - defer func() { config.Config.Redis.Enabled = origRedisEnabled }() + origAddr := config.Config.App.Addr + config.Config.App.Addr = "127.0.0.1:0" + defer func() { + config.Config.Redis.Enabled = origRedisEnabled + config.Config.App.Addr = origAddr + }() // ══════════════════════════════════════════════════════════════════════════ // 场景 1: 拔出 Redis (Zero-Redis Monolith 模式) @@ -174,7 +179,7 @@ func TestRedisPluggability_Simulation(t *testing.T) { require.Eventually(t, func() bool { return asynqTaskExecuted.Load() >= 1 - }, 5*time.Second, 100*time.Millisecond, "Asynq Worker 应从 Redis 队列中成功消费并执行任务") + }, 10*time.Second, 100*time.Millisecond, "Asynq Worker 应从 Redis 队列中成功消费并执行任务") // 6. 优雅关闭 stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second) diff --git a/backend/core/app_test.go b/backend/core/app_test.go index 484927c8..71aca568 100644 --- a/backend/core/app_test.go +++ b/backend/core/app_test.go @@ -406,10 +406,10 @@ func TestAppRunContextCancellation(t *testing.T) { errCh <- app.Run(ctx) }() - // Wait briefly then cancel - time.Sleep(50 * time.Millisecond) - assert.True(t, app.IsRunning()) - assert.True(t, d.isStarted()) + // Wait for app and driver to become ready + assert.Eventually(t, func() bool { + return app.IsRunning() && d.isStarted() + }, 2*time.Second, 10*time.Millisecond) cancel() diff --git a/backend/core/context.go b/backend/core/context.go index 772665d7..d7e69b87 100644 --- a/backend/core/context.go +++ b/backend/core/context.go @@ -162,12 +162,12 @@ func (c *Context) ForkWithContext(base context.Context) *Context { cancel: cancel, parent: c, container: NewContainer(c.container), - events: c.Events(), - router: c.Router(), - migrations: c.Migrations(), - tasks: c.Tasks(), - schedules: c.Schedules(), - settings: c.Settings(), + events: c.events, + router: c.router, + migrations: c.migrations, + tasks: c.tasks, + schedules: c.schedules, + settings: c.settings, values: make(map[any]any), } @@ -354,20 +354,22 @@ func (c *Context) RegisterDriver(d Driver) error { // Drivers returns a copy of all drivers registered on this Context. func (c *Context) Drivers() []Driver { - c.mu.RLock() - defer c.mu.RUnlock() + root := c.Root() + root.mu.RLock() + defer root.mu.RUnlock() - result := make([]Driver, len(c.drivers)) - copy(result, c.drivers) + result := make([]Driver, len(root.drivers)) + copy(result, root.drivers) return result } // Driver looks up a registered driver by its driver type. func (c *Context) Driver(driverType DriverType) (Driver, bool) { - c.mu.RLock() - defer c.mu.RUnlock() + root := c.Root() + root.mu.RLock() + defer root.mu.RUnlock() - for _, d := range c.drivers { + for _, d := range root.drivers { if d.Type() == driverType { return d, true } diff --git a/backend/plugins/drivers/driver_asynq_worker/executor.go b/backend/plugins/drivers/driver_asynq_worker/executor.go index 40c3d5e2..c1776aea 100644 --- a/backend/plugins/drivers/driver_asynq_worker/executor.go +++ b/backend/plugins/drivers/driver_asynq_worker/executor.go @@ -154,7 +154,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere // 入队 Asynq taskInfo := asynq.NewTask(meta.AsynqTask, injectTaskTraceContext(ctx, payload)) - if _, err := AsynqClient.Enqueue( + if _, err := GetAsynqClient().Enqueue( taskInfo, asynq.TaskID(taskID), asynq.MaxRetry(meta.MaxRetry), @@ -220,7 +220,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) { // 入队 Asynq taskInfo := asynq.NewTask(execution.TaskType, injectTaskTraceContext(ctx, []byte(execution.Payload))) - if _, err := AsynqClient.Enqueue( + if _, err := GetAsynqClient().Enqueue( taskInfo, asynq.TaskID(newTaskID), asynq.MaxRetry(execution.MaxRetry), diff --git a/backend/plugins/drivers/driver_asynq_worker/plugin.go b/backend/plugins/drivers/driver_asynq_worker/plugin.go index 921743a2..5dc3dc07 100644 --- a/backend/plugins/drivers/driver_asynq_worker/plugin.go +++ b/backend/plugins/drivers/driver_asynq_worker/plugin.go @@ -13,10 +13,10 @@ import ( "time" "github.com/hibiken/asynq" + "github.com/redis/go-redis/v9" "Wavelet/core" "Wavelet/core/contracts" - "Wavelet/pkg/config" ) const ( @@ -90,7 +90,6 @@ type Plugin struct { // New creates a new Asynq Worker driver plugin. func New(opts ...Option) *Plugin { p := &Plugin{ - redisOpt: RedisOpt, concurrency: defaultConcurrency, shutdownTimeout: defaultShutdownTimeout, queues: map[string]int{"default": 1}, @@ -139,6 +138,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { ctx.OnDispose(func() error { SetActiveTaskExtension(nil) + SetRedisClient(nil) + ResetAsynqClient() shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout) defer cancel() return p.Stop(shutdownCtx) @@ -173,29 +174,21 @@ func (p *Plugin) Start(_ context.Context) error { } } - for _, taskName := range GetRegisteredAsynqTasks() { - mux.HandleFunc(taskName, ProcessTask) + opt := p.redisOpt + if opt == nil { + opt = NewRedisConnOpt() + } + RedisOpt = opt + + if getRedisClient() == nil { + if mk, ok := opt.(interface{ MakeRedisClient() interface{} }); ok { + if client, ok := mk.MakeRedisClient().(redis.UniversalClient); ok { + SetRedisClient(client) + } + } } if p.server == nil { - opt := p.redisOpt - if opt == nil { - if RedisOpt != nil { - opt = RedisOpt - } else { - redisCfg := config.Config.Redis - addr := "127.0.0.1:6379" - if len(redisCfg.Addrs) > 0 && redisCfg.Addrs[0] != "" { - addr = redisCfg.Addrs[0] - } - opt = asynq.RedisClientOpt{ - Addr: addr, - Username: redisCfg.Username, - Password: redisCfg.Password, - DB: redisCfg.DB, - } - } - } p.server = asynq.NewServer( opt, asynq.Config{ @@ -276,15 +269,17 @@ func toAsynqHandler(h any) (asynq.Handler, error) { return asynq.HandlerFunc(fn), nil case func(context.Context, []byte) error: return asynq.HandlerFunc(func(c context.Context, t *asynq.Task) error { - return fn(c, t.Payload()) + c, payload, _ := extractTaskTraceContext(c, t.Payload()) + return fn(c, payload) }), nil case func(context.Context) error: return asynq.HandlerFunc(func(c context.Context, _ *asynq.Task) error { return fn(c) }), nil case func([]byte) error: - return asynq.HandlerFunc(func(_ context.Context, t *asynq.Task) error { - return fn(t.Payload()) + return asynq.HandlerFunc(func(c context.Context, t *asynq.Task) error { + _, payload, _ := extractTaskTraceContext(c, t.Payload()) + return fn(payload) }), nil case func() error: return asynq.HandlerFunc(func(_ context.Context, _ *asynq.Task) error { diff --git a/backend/plugins/drivers/driver_asynq_worker/utils.go b/backend/plugins/drivers/driver_asynq_worker/utils.go index 25121401..fa275d6c 100644 --- a/backend/plugins/drivers/driver_asynq_worker/utils.go +++ b/backend/plugins/drivers/driver_asynq_worker/utils.go @@ -4,6 +4,8 @@ package driver_asynq_worker import ( + "sync" + "github.com/hibiken/asynq" "github.com/redis/go-redis/v9" "github.com/redis/go-redis/v9/maintnotifications" @@ -54,9 +56,37 @@ var RedisOpt asynq.RedisConnOpt // AsynqClient asynq 客户端,用于任务入队 var AsynqClient *asynq.Client -func init() { - RedisOpt = NewRedisConnOpt() - AsynqClient = asynq.NewClient(RedisOpt) +var asynqClientMu sync.RWMutex + +// GetAsynqClient returns the active Asynq client, dynamically creating or refreshing it if needed. +func GetAsynqClient() *asynq.Client { + asynqClientMu.RLock() + if AsynqClient != nil { + asynqClientMu.RUnlock() + return AsynqClient + } + asynqClientMu.RUnlock() + + asynqClientMu.Lock() + defer asynqClientMu.Unlock() + if AsynqClient != nil { + return AsynqClient + } + + opt := NewRedisConnOpt() + RedisOpt = opt + AsynqClient = asynq.NewClient(opt) + return AsynqClient +} + +// ResetAsynqClient resets the Asynq client so it will be re-created with current config. +func ResetAsynqClient() { + asynqClientMu.Lock() + defer asynqClientMu.Unlock() + if AsynqClient != nil { + _ = AsynqClient.Close() + AsynqClient = nil + } } // NewRedisConnOpt 根据配置返回对应的 asynq Redis 连接选项