fix(core): delegate driver and task extension lookups to root context and improve test reliability

This commit is contained in:
ryan
2026-08-28 16:09:18 +08:00
parent 7fca53823a
commit 3a61c38dc5
8 changed files with 89 additions and 52 deletions
+1
View File
@@ -65,3 +65,4 @@ s3_cache
/.superpowers/ /.superpowers/
/backend/plugins/domain/upload/filesrv/uploads/ /backend/plugins/domain/upload/filesrv/uploads/
/backend/plugins/domain/upload/task/uploads/ /backend/plugins/domain/upload/task/uploads/
/backend/data/
+7 -3
View File
@@ -113,12 +113,16 @@ type sharedStore struct {
func (s *sharedStore) Tablename() string { return "w_schema_versions" } func (s *sharedStore) Tablename() string { return "w_schema_versions" }
func (s *sharedStore) CreateVersionTable(ctx context.Context, db goosedb.DBTxConn) error { 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, plugin_id VARCHAR(64) NOT NULL,
version_id BIGINT 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) PRIMARY KEY (plugin_id, version_id)
)`) )`, timeType))
return err return err
} }
+7 -2
View File
@@ -21,7 +21,12 @@ import (
func TestRedisPluggability_Simulation(t *testing.T) { func TestRedisPluggability_Simulation(t *testing.T) {
origRedisEnabled := config.Config.Redis.Enabled 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 模式) // 场景 1: 拔出 Redis (Zero-Redis Monolith 模式)
@@ -174,7 +179,7 @@ func TestRedisPluggability_Simulation(t *testing.T) {
require.Eventually(t, func() bool { require.Eventually(t, func() bool {
return asynqTaskExecuted.Load() >= 1 return asynqTaskExecuted.Load() >= 1
}, 5*time.Second, 100*time.Millisecond, "Asynq Worker 应从 Redis 队列中成功消费并执行任务") }, 10*time.Second, 100*time.Millisecond, "Asynq Worker 应从 Redis 队列中成功消费并执行任务")
// 6. 优雅关闭 // 6. 优雅关闭
stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second) stopCtx, stopCancel := context.WithTimeout(context.Background(), 5*time.Second)
+4 -4
View File
@@ -406,10 +406,10 @@ func TestAppRunContextCancellation(t *testing.T) {
errCh <- app.Run(ctx) errCh <- app.Run(ctx)
}() }()
// Wait briefly then cancel // Wait for app and driver to become ready
time.Sleep(50 * time.Millisecond) assert.Eventually(t, func() bool {
assert.True(t, app.IsRunning()) return app.IsRunning() && d.isStarted()
assert.True(t, d.isStarted()) }, 2*time.Second, 10*time.Millisecond)
cancel() cancel()
+15 -13
View File
@@ -162,12 +162,12 @@ func (c *Context) ForkWithContext(base context.Context) *Context {
cancel: cancel, cancel: cancel,
parent: c, parent: c,
container: NewContainer(c.container), container: NewContainer(c.container),
events: c.Events(), events: c.events,
router: c.Router(), router: c.router,
migrations: c.Migrations(), migrations: c.migrations,
tasks: c.Tasks(), tasks: c.tasks,
schedules: c.Schedules(), schedules: c.schedules,
settings: c.Settings(), settings: c.settings,
values: make(map[any]any), 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. // Drivers returns a copy of all drivers registered on this Context.
func (c *Context) Drivers() []Driver { func (c *Context) Drivers() []Driver {
c.mu.RLock() root := c.Root()
defer c.mu.RUnlock() root.mu.RLock()
defer root.mu.RUnlock()
result := make([]Driver, len(c.drivers)) result := make([]Driver, len(root.drivers))
copy(result, c.drivers) copy(result, root.drivers)
return result return result
} }
// Driver looks up a registered driver by its driver type. // Driver looks up a registered driver by its driver type.
func (c *Context) Driver(driverType DriverType) (Driver, bool) { func (c *Context) Driver(driverType DriverType) (Driver, bool) {
c.mu.RLock() root := c.Root()
defer c.mu.RUnlock() root.mu.RLock()
defer root.mu.RUnlock()
for _, d := range c.drivers { for _, d := range root.drivers {
if d.Type() == driverType { if d.Type() == driverType {
return d, true return d, true
} }
@@ -154,7 +154,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere
// 入队 Asynq // 入队 Asynq
taskInfo := asynq.NewTask(meta.AsynqTask, injectTaskTraceContext(ctx, payload)) taskInfo := asynq.NewTask(meta.AsynqTask, injectTaskTraceContext(ctx, payload))
if _, err := AsynqClient.Enqueue( if _, err := GetAsynqClient().Enqueue(
taskInfo, taskInfo,
asynq.TaskID(taskID), asynq.TaskID(taskID),
asynq.MaxRetry(meta.MaxRetry), asynq.MaxRetry(meta.MaxRetry),
@@ -220,7 +220,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) {
// 入队 Asynq // 入队 Asynq
taskInfo := asynq.NewTask(execution.TaskType, injectTaskTraceContext(ctx, []byte(execution.Payload))) taskInfo := asynq.NewTask(execution.TaskType, injectTaskTraceContext(ctx, []byte(execution.Payload)))
if _, err := AsynqClient.Enqueue( if _, err := GetAsynqClient().Enqueue(
taskInfo, taskInfo,
asynq.TaskID(newTaskID), asynq.TaskID(newTaskID),
asynq.MaxRetry(execution.MaxRetry), asynq.MaxRetry(execution.MaxRetry),
@@ -13,10 +13,10 @@ import (
"time" "time"
"github.com/hibiken/asynq" "github.com/hibiken/asynq"
"github.com/redis/go-redis/v9"
"Wavelet/core" "Wavelet/core"
"Wavelet/core/contracts" "Wavelet/core/contracts"
"Wavelet/pkg/config"
) )
const ( const (
@@ -90,7 +90,6 @@ type Plugin struct {
// New creates a new Asynq Worker driver plugin. // New creates a new Asynq Worker driver plugin.
func New(opts ...Option) *Plugin { func New(opts ...Option) *Plugin {
p := &Plugin{ p := &Plugin{
redisOpt: RedisOpt,
concurrency: defaultConcurrency, concurrency: defaultConcurrency,
shutdownTimeout: defaultShutdownTimeout, shutdownTimeout: defaultShutdownTimeout,
queues: map[string]int{"default": 1}, queues: map[string]int{"default": 1},
@@ -139,6 +138,8 @@ func (p *Plugin) Apply(ctx *core.Context) error {
ctx.OnDispose(func() error { ctx.OnDispose(func() error {
SetActiveTaskExtension(nil) SetActiveTaskExtension(nil)
SetRedisClient(nil)
ResetAsynqClient()
shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout) shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout)
defer cancel() defer cancel()
return p.Stop(shutdownCtx) return p.Stop(shutdownCtx)
@@ -173,29 +174,21 @@ func (p *Plugin) Start(_ context.Context) error {
} }
} }
for _, taskName := range GetRegisteredAsynqTasks() { opt := p.redisOpt
mux.HandleFunc(taskName, ProcessTask) 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 { 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( p.server = asynq.NewServer(
opt, opt,
asynq.Config{ asynq.Config{
@@ -276,15 +269,17 @@ func toAsynqHandler(h any) (asynq.Handler, error) {
return asynq.HandlerFunc(fn), nil return asynq.HandlerFunc(fn), nil
case func(context.Context, []byte) error: case func(context.Context, []byte) error:
return asynq.HandlerFunc(func(c context.Context, t *asynq.Task) 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 }), nil
case func(context.Context) error: case func(context.Context) error:
return asynq.HandlerFunc(func(c context.Context, _ *asynq.Task) error { return asynq.HandlerFunc(func(c context.Context, _ *asynq.Task) error {
return fn(c) return fn(c)
}), nil }), nil
case func([]byte) error: case func([]byte) error:
return asynq.HandlerFunc(func(_ context.Context, t *asynq.Task) error { return asynq.HandlerFunc(func(c context.Context, t *asynq.Task) error {
return fn(t.Payload()) _, payload, _ := extractTaskTraceContext(c, t.Payload())
return fn(payload)
}), nil }), nil
case func() error: case func() error:
return asynq.HandlerFunc(func(_ context.Context, _ *asynq.Task) error { return asynq.HandlerFunc(func(_ context.Context, _ *asynq.Task) error {
@@ -4,6 +4,8 @@
package driver_asynq_worker package driver_asynq_worker
import ( import (
"sync"
"github.com/hibiken/asynq" "github.com/hibiken/asynq"
"github.com/redis/go-redis/v9" "github.com/redis/go-redis/v9"
"github.com/redis/go-redis/v9/maintnotifications" "github.com/redis/go-redis/v9/maintnotifications"
@@ -54,9 +56,37 @@ var RedisOpt asynq.RedisConnOpt
// AsynqClient asynq 客户端,用于任务入队 // AsynqClient asynq 客户端,用于任务入队
var AsynqClient *asynq.Client var AsynqClient *asynq.Client
func init() { var asynqClientMu sync.RWMutex
RedisOpt = NewRedisConnOpt()
AsynqClient = asynq.NewClient(RedisOpt) // 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 连接选项 // NewRedisConnOpt 根据配置返回对应的 asynq Redis 连接选项