topfans/backend/pkg/mq/asynq/adapter.go

319 lines
8.2 KiB
Go

// 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
}