319 lines
8.2 KiB
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
|
|
}
|