// Package asynq 是 adapter.Adapter 的 Asynq 实现,只提供 Task 原语。 // // 事件原语由 sibling package streams 实现,合成由 pkg/mq/mq.go 完成。 package asynq import ( "context" "encoding/json" "errors" "fmt" "sync" "time" "github.com/hibiken/asynq" "github.com/topfans/backend/pkg/logger" "github.com/topfans/backend/pkg/mq/adapter" "go.uber.org/zap" ) // Adapter 同时实现 TaskProducer 和 TaskConsumer 两个接口, // 它们共用一个 redis 连接配置的 asynq.Client。 type Adapter struct { cfg Config client *asynq.Client mu sync.Mutex server *asynq.Server mux *asynq.ServeMux running bool // cron 登记的 spec → taskType,Run 时注册到 scheduler cronSpecs []cronEntry } type cronEntry struct { spec string taskType string payload map[string]any } // New 构造 asynq adapter,建立 redis client 连接。 func New(cfg Config) (*Adapter, error) { if cfg.RedisAddr == "" { return nil, errors.New("asynq: RedisAddr is required") } if cfg.Concurrency <= 0 { cfg.Concurrency = 10 } if len(cfg.Queues) == 0 { cfg.Queues = map[string]int{"default": 1} } redisOpt := asynq.RedisClientOpt{ Addr: cfg.RedisAddr, DB: cfg.RedisDB, Password: cfg.Password, } client := asynq.NewClient(redisOpt) // 验证连接 if err := client.Ping(); err != nil { return nil, fmt.Errorf("asynq: redis ping failed: %w", err) } return &Adapter{ cfg: cfg, client: client, mux: asynq.NewServeMux(), }, nil } // Close 关闭 server 和 client。 func (a *Adapter) Close() error { a.mu.Lock() defer a.mu.Unlock() if a.server != nil { a.server.Stop() a.server = nil } if a.client != nil { return a.client.Close() } return nil } // ============================================================= // Producer 实现 (adapter.TaskProducer) // ============================================================= // Enqueue 实现 adapter.TaskProducer.Enqueue。 func (a *Adapter) Enqueue(ctx context.Context, t adapter.Task) (string, error) { return a.enqueue(ctx, t) } // EnqueueAt 实现 adapter.TaskProducer.EnqueueAt。 func (a *Adapter) EnqueueAt(ctx context.Context, t adapter.Task, processAt time.Time) (string, error) { // 业务侧 ProcessAt 字段也可能是零值,优先级:函数入参 > t.ProcessAt if !processAt.IsZero() { t.ProcessAt = processAt } return a.enqueue(ctx, t) } // EnqueueUnique 实现 adapter.TaskProducer.EnqueueUnique。 func (a *Adapter) EnqueueUnique(ctx context.Context, t adapter.Task, ttl time.Duration) (string, error) { t.UniqueTTL = ttl return a.enqueue(ctx, t) } // GetInfo 占位实现 — Asynq v0.26 不在 public API 中提供 GetInspector. // 业务侧需要的话通过 asynqmon dashboard 查询,后续可以替换为 inspector 查询。 func (a *Adapter) GetInfo(ctx context.Context, queue, taskID string) (*adapter.TaskInfo, error) { return nil, adapter.ErrTaskNotFound } // Delete 占位实现 — 同上 Asynq v0.26 不在 public API 提供。 func (a *Adapter) Delete(ctx context.Context, queue, taskID string) error { return errors.New("mq: asynq adapter does not support Delete in current version") } // enqueue 内部统一入队方法,负责 Task → asynq.TaskOptions 映射。 func (a *Adapter) enqueue(ctx context.Context, t adapter.Task) (string, error) { if err := t.IsValid(); err != nil { return "", err } mp := marshalPayload(t.Payload) payloadBytes, err := mp.Value() if err != nil { return "", fmt.Errorf("asynq: marshal payload: %w", err) } task := asynq.NewTask(t.Type, payloadBytes, asynq.MaxRetry(resolveMaxRetry(t.MaxRetry))) opts := []asynq.Option{ asynq.Queue(resolveQueue(t.Queue)), asynq.Timeout(resolveTimeout(t.Timeout)), } if !t.ProcessAt.IsZero() { opts = append(opts, asynq.ProcessAt(t.ProcessAt)) } if t.UniqueTTL > 0 { opts = append(opts, asynq.Unique(t.UniqueTTL)) if t.UniqueKey != "" { opts = append(opts, asynq.TaskID(t.UniqueKey)) } } info, err := a.client.Enqueue(task, opts...) if err != nil { if errors.Is(err, asynq.ErrDuplicateTask) { return "", fmt.Errorf("mq: duplicate task: %w", err) } logger.Logger.Warn("asynq enqueue failed", zap.String("type", t.Type), zap.Error(err)) return "", fmt.Errorf("asynq: enqueue: %w", err) } return info.ID, nil } // ============================================================= // Consumer 实现 (adapter.TaskConsumer) // ============================================================= // RegisterTask 实现 adapter.TaskConsumer.RegisterTask。 func (a *Adapter) RegisterTask(taskType string, handler adapter.TaskHandler, opts adapter.TaskRegisterOptions) error { a.mu.Lock() defer a.mu.Unlock() wrapped := a.wrapHandler(taskType, handler) a.mux.HandleFunc(taskType, wrapped) return nil } // RegisterCron 实现 adapter.TaskConsumer.RegisterCron。 func (a *Adapter) RegisterCron(spec string, taskType string, payload map[string]any) error { a.mu.Lock() defer a.mu.Unlock() a.cronSpecs = append(a.cronSpecs, cronEntry{spec, taskType, payload}) return nil } // Run 启动 server + cron scheduler。 func (a *Adapter) Run(ctx context.Context) error { a.mu.Lock() if a.running { a.mu.Unlock() return adapter.ErrConsumerAlreadyRunning } redisOpt := asynq.RedisClientOpt{ Addr: a.cfg.RedisAddr, DB: a.cfg.RedisDB, Password: a.cfg.Password, } // 注册 cron specs scheduler := asynq.NewScheduler(redisOpt, &asynq.SchedulerOpts{ Location: time.Local, }) for _, c := range a.cronSpecs { mp := marshalPayload(c.payload) payloadBytes, err := mp.Value() if err != nil { a.mu.Unlock() return fmt.Errorf("asynq: marshal cron payload: %w", err) } task := asynq.NewTask(c.taskType, payloadBytes) if _, err := scheduler.Register(c.spec, task); err != nil { a.mu.Unlock() return fmt.Errorf("asynq: register cron %q: %w", c.spec, err) } } srv := asynq.NewServer(redisOpt, asynq.Config{ Concurrency: a.cfg.Concurrency, Queues: a.cfg.Queues, ErrorHandler: asynq.ErrorHandlerFunc(func(ctx context.Context, task *asynq.Task, err error) { taskID := "unknown" if rw := task.ResultWriter(); rw != nil { taskID = rw.TaskID() } logger.Logger.Error("asynq task error", zap.String("type", task.Type()), zap.String("id", taskID), zap.Error(err)) }), }) a.server = srv a.running = true a.mu.Unlock() // 启动两个并行 goroutine:server + scheduler,任何 ctx done 即退出 errCh := make(chan error, 2) go func() { if err := scheduler.Run(); err != nil { errCh <- fmt.Errorf("scheduler: %w", err) } }() go func() { if err := srv.Run(a.mux); err != nil { errCh <- fmt.Errorf("server: %w", err) } }() select { case <-ctx.Done(): srv.Shutdown() scheduler.Shutdown() return ctx.Err() case err := <-errCh: srv.Shutdown() scheduler.Shutdown() return err } } // Stop 优雅停机。 func (a *Adapter) Stop() error { a.mu.Lock() srv := a.server a.server = nil a.running = false a.mu.Unlock() if srv != nil { srv.Shutdown() } return nil } // wrapHandler 把业务侧 adapter.TaskHandler 包装成 asynq.HandlerFunc。 // 内置 panic recovery — handler panic 不会导致整个 server 崩。 func (a *Adapter) wrapHandler(taskType string, handler adapter.TaskHandler) func(context.Context, *asynq.Task) error { return func(ctx context.Context, t *asynq.Task) error { defer func() { if r := recover(); r != nil { logger.Logger.Error("asynq handler panic recovered", zap.String("type", taskType), zap.String("task_id", t.ResultWriter().TaskID()), zap.Any("panic", r)) } }() atask := &adapter.Task{ Type: t.Type(), Payload: map[string]any{}, } p := t.Payload() if len(p) > 0 { m := map[string]any{} if err := json.Unmarshal(p, &m); err != nil { return fmt.Errorf("asynq: unmarshal payload %q: %w", taskType, err) } atask.Payload = m } return handler(ctx, atask) } } // ============================================================= // helpers // ============================================================= func resolveQueue(q string) string { if q == "" { return "default" } return q } func resolveTimeout(d time.Duration) time.Duration { return d // asynq 表示 0 = 无超时 } func resolveMaxRetry(n int) int { if n < 0 { return 25 // asynq 默认 } return n }