diff --git a/cmd/sweeper/main.go b/cmd/sweeper/main.go new file mode 100644 index 0000000..7825331 --- /dev/null +++ b/cmd/sweeper/main.go @@ -0,0 +1,68 @@ +package main + +import ( + "context" + "log/slog" + "os" + "os/signal" + "syscall" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + goredis "github.com/redis/go-redis/v9" + + "github.com/smallchungus/disttaskqueue/internal/queue" + "github.com/smallchungus/disttaskqueue/internal/store" + "github.com/smallchungus/disttaskqueue/internal/sweeper" +) + +func main() { + logger := slog.New(slog.NewJSONHandler(os.Stdout, nil)) + slog.SetDefault(logger) + + dsn := envOr("DATABASE_URL", "postgres://dtq:dtq@localhost:5432/dtq?sslmode=disable") + redisURL := envOr("REDIS_URL", "redis://localhost:6379/0") + + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + if err := store.Migrate(ctx, dsn); err != nil { + slog.Error("migrate", "err", err) + os.Exit(1) + } + + pool, err := pgxpool.New(ctx, dsn) + if err != nil { + slog.Error("pg connect", "err", err) + os.Exit(1) + } + defer pool.Close() + + opts, err := goredis.ParseURL(redisURL) + if err != nil { + slog.Error("redis url", "err", err) + os.Exit(1) + } + redis := goredis.NewClient(opts) + defer func() { _ = redis.Close() }() + + sw := sweeper.New(sweeper.Config{ + Store: store.New(pool), + Queue: queue.New(redis), + Interval: 5 * time.Second, + }) + + slog.Info("sweeper starting") + if err := sw.Run(ctx); err != nil { + slog.Error("sweeper run", "err", err) + os.Exit(1) + } + slog.Info("sweeper stopped") +} + +func envOr(k, def string) string { + if v := os.Getenv(k); v != "" { + return v + } + return def +} diff --git a/internal/queue/queue.go b/internal/queue/queue.go index 831c0f8..48b8ca5 100644 --- a/internal/queue/queue.go +++ b/internal/queue/queue.go @@ -56,6 +56,14 @@ func (q *Queue) Depth(ctx context.Context, stage string) (int64, error) { // Client exposes the underlying Redis client. Test-only. func (q *Queue) Client() *goredis.Client { return q.cli } +func (q *Queue) IsWorkerAlive(ctx context.Context, workerID string) (bool, error) { + n, err := q.cli.Exists(ctx, "heartbeat:"+workerID).Result() + if err != nil { + return false, fmt.Errorf("exists: %w", err) + } + return n > 0, nil +} + func (q *Queue) Heartbeat(ctx context.Context, workerID string, ttl time.Duration) error { now := strconv.FormatInt(time.Now().Unix(), 10) if err := q.cli.Set(ctx, "heartbeat:"+workerID, now, ttl).Err(); err != nil { diff --git a/internal/queue/queue_integration_test.go b/internal/queue/queue_integration_test.go index a5eab07..0168741 100644 --- a/internal/queue/queue_integration_test.go +++ b/internal/queue/queue_integration_test.go @@ -82,6 +82,32 @@ func TestBlockingPop_ReturnsErrEmptyOnTimeout(t *testing.T) { } } +func TestIsWorkerAlive_TrueWhenHeartbeatSet(t *testing.T) { + q := newQueue(t) + ctx := context.Background() + if err := q.Heartbeat(ctx, "w1", 5*time.Second); err != nil { + t.Fatal(err) + } + alive, err := q.IsWorkerAlive(ctx, "w1") + if err != nil { + t.Fatal(err) + } + if !alive { + t.Fatal("expected alive=true") + } +} + +func TestIsWorkerAlive_FalseWhenHeartbeatMissing(t *testing.T) { + q := newQueue(t) + alive, err := q.IsWorkerAlive(context.Background(), "never-lived") + if err != nil { + t.Fatal(err) + } + if alive { + t.Fatal("expected alive=false") + } +} + func TestHeartbeat_SetsKeyWithTTL(t *testing.T) { q := newQueue(t) cli := q.Client() diff --git a/internal/store/store.go b/internal/store/store.go index 024e7b0..9f1b9fa 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -101,6 +101,49 @@ func (s *Store) MarkFailed(ctx context.Context, id uuid.UUID, errMsg string, nex return nil } +func (s *Store) ListRunningJobs(ctx context.Context) ([]Job, error) { + const q = ` + SELECT id, stage, status, payload, worker_id, attempts, max_attempts, + last_error, next_run_at, claimed_at, completed_at, created_at, updated_at + FROM pipeline_jobs WHERE status = $1` + rows, err := s.pool.Query(ctx, q, StatusRunning) + if err != nil { + return nil, fmt.Errorf("list running: %w", err) + } + defer rows.Close() + return scanJobs(rows) +} + +func (s *Store) ListReadyRetryJobs(ctx context.Context) ([]Job, error) { + const q = ` + SELECT id, stage, status, payload, worker_id, attempts, max_attempts, + last_error, next_run_at, claimed_at, completed_at, created_at, updated_at + FROM pipeline_jobs + WHERE status = $1 AND last_error IS NOT NULL AND next_run_at <= now()` + rows, err := s.pool.Query(ctx, q, StatusQueued) + if err != nil { + return nil, fmt.Errorf("list ready retries: %w", err) + } + defer rows.Close() + return scanJobs(rows) +} + +func scanJobs(rows pgx.Rows) ([]Job, error) { + var out []Job + for rows.Next() { + var j Job + if err := rows.Scan(&j.ID, &j.Stage, &j.Status, &j.Payload, &j.WorkerID, &j.Attempts, + &j.MaxAttempts, &j.LastError, &j.NextRunAt, &j.ClaimedAt, &j.CompletedAt, &j.CreatedAt, &j.UpdatedAt); err != nil { + return nil, fmt.Errorf("scan: %w", err) + } + out = append(out, j) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("rows: %w", err) + } + return out, nil +} + func (s *Store) GetJob(ctx context.Context, id uuid.UUID) (Job, error) { const q = ` SELECT id, stage, status, payload, worker_id, attempts, max_attempts, diff --git a/internal/store/store_integration_test.go b/internal/store/store_integration_test.go index 7587aef..43b08a8 100644 --- a/internal/store/store_integration_test.go +++ b/internal/store/store_integration_test.go @@ -145,6 +145,67 @@ func TestMarkFailed_RequeuesWithIncrementedAttempts(t *testing.T) { } } +func TestListRunningJobs_ReturnsOnlyRunning(t *testing.T) { + s := newStore(t) + ctx := context.Background() + + jQueued, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + jRunning, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, jRunning.ID, "w1") + jDone, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, jDone.ID, "w2") + _ = s.MarkDone(ctx, jDone.ID) + + jobs, err := s.ListRunningJobs(ctx) + if err != nil { + t.Fatal(err) + } + if len(jobs) != 1 { + t.Fatalf("got %d jobs, want 1", len(jobs)) + } + if jobs[0].ID != jRunning.ID { + t.Fatalf("got id %s, want %s", jobs[0].ID, jRunning.ID) + } + _ = jQueued +} + +func TestListReadyRetryJobs_ReturnsQueuedWithLastErrorAndDueNextRun(t *testing.T) { + s := newStore(t) + ctx := context.Background() + + j1, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, j1.ID, "w1") + past := time.Now().Add(-1 * time.Minute) + _ = s.MarkFailed(ctx, j1.ID, "boom", past) + + j2, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, j2.ID, "w1") + future := time.Now().Add(5 * time.Minute) + _ = s.MarkFailed(ctx, j2.ID, "boom", future) + + j3, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + + jobs, err := s.ListReadyRetryJobs(ctx) + if err != nil { + t.Fatal(err) + } + if len(jobs) != 1 { + t.Fatalf("got %d, want 1; got IDs: %v", len(jobs), jobIDs(jobs)) + } + if jobs[0].ID != j1.ID { + t.Fatalf("got %s, want %s", jobs[0].ID, j1.ID) + } + _, _ = j2, j3 +} + +func jobIDs(js []store.Job) []string { + out := make([]string, len(js)) + for i, j := range js { + out[i] = j.ID.String() + } + return out +} + func TestMarkFailed_MarksDeadAtMaxAttempts(t *testing.T) { s := newStore(t) ctx := context.Background() diff --git a/internal/sweeper/sweeper.go b/internal/sweeper/sweeper.go new file mode 100644 index 0000000..e4a5d04 --- /dev/null +++ b/internal/sweeper/sweeper.go @@ -0,0 +1,95 @@ +package sweeper + +import ( + "context" + "fmt" + "log/slog" + "time" + + "github.com/smallchungus/disttaskqueue/internal/queue" + "github.com/smallchungus/disttaskqueue/internal/store" +) + +type Config struct { + Store *store.Store + Queue *queue.Queue + Interval time.Duration +} + +type Sweeper struct { + cfg Config +} + +func New(cfg Config) *Sweeper { + if cfg.Interval == 0 { + cfg.Interval = 5 * time.Second + } + return &Sweeper{cfg: cfg} +} + +func (s *Sweeper) Run(ctx context.Context) error { + tick := time.NewTicker(s.cfg.Interval) + defer tick.Stop() + + for { + select { + case <-ctx.Done(): + return nil + case <-tick.C: + if err := s.SweepOnce(ctx); err != nil { + slog.Error("sweep failed", "err", err) + } + } + } +} + +func (s *Sweeper) SweepOnce(ctx context.Context) error { + if err := s.reviveOrphans(ctx); err != nil { + return fmt.Errorf("revive orphans: %w", err) + } + if err := s.promoteDelayed(ctx); err != nil { + return fmt.Errorf("promote delayed: %w", err) + } + return nil +} + +func (s *Sweeper) reviveOrphans(ctx context.Context) error { + running, err := s.cfg.Store.ListRunningJobs(ctx) + if err != nil { + return err + } + for _, j := range running { + if j.WorkerID == nil { + continue + } + alive, err := s.cfg.Queue.IsWorkerAlive(ctx, *j.WorkerID) + if err != nil { + slog.Warn("heartbeat check failed", "job_id", j.ID, "err", err) + continue + } + if alive { + continue + } + if err := s.cfg.Store.MarkFailed(ctx, j.ID, "worker died", time.Now()); err != nil { + slog.Warn("revive failed", "job_id", j.ID, "err", err) + continue + } + slog.Info("revived orphan", "job_id", j.ID, "worker_id", *j.WorkerID) + } + return nil +} + +func (s *Sweeper) promoteDelayed(ctx context.Context) error { + ready, err := s.cfg.Store.ListReadyRetryJobs(ctx) + if err != nil { + return err + } + for _, j := range ready { + if err := s.cfg.Queue.Push(ctx, j.Stage, j.ID.String()); err != nil { + slog.Warn("promote failed", "job_id", j.ID, "err", err) + continue + } + slog.Info("promoted retry", "job_id", j.ID, "stage", j.Stage) + } + return nil +} diff --git a/internal/sweeper/sweeper_integration_test.go b/internal/sweeper/sweeper_integration_test.go new file mode 100644 index 0000000..fdd2d04 --- /dev/null +++ b/internal/sweeper/sweeper_integration_test.go @@ -0,0 +1,121 @@ +//go:build integration + +package sweeper_test + +import ( + "context" + "testing" + "time" + + "github.com/smallchungus/disttaskqueue/internal/queue" + "github.com/smallchungus/disttaskqueue/internal/store" + "github.com/smallchungus/disttaskqueue/internal/sweeper" + "github.com/smallchungus/disttaskqueue/internal/testutil" +) + +func setup(t *testing.T) (*store.Store, *queue.Queue, *sweeper.Sweeper) { + t.Helper() + pool := testutil.StartPostgres(t) + if err := store.Migrate(context.Background(), pool.Config().ConnString()); err != nil { + t.Fatal(err) + } + s := store.New(pool) + q := queue.New(testutil.StartRedis(t)) + sw := sweeper.New(sweeper.Config{Store: s, Queue: q}) + return s, q, sw +} + +func TestSweepOnce_RequeuesJobWithDeadWorker(t *testing.T) { + s, q, sw := setup(t) + ctx := context.Background() + + if err := q.Heartbeat(ctx, "w1", 500*time.Millisecond); err != nil { + t.Fatal(err) + } + job, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + if err := s.ClaimJob(ctx, job.ID, "w1"); err != nil { + t.Fatal(err) + } + time.Sleep(700 * time.Millisecond) + + if err := sw.SweepOnce(ctx); err != nil { + t.Fatalf("sweep: %v", err) + } + + got, _ := s.GetJob(ctx, job.ID) + if got.Status != store.StatusQueued { + t.Fatalf("status: %s, want queued (revived)", got.Status) + } + if got.Attempts != 1 { + t.Fatalf("attempts: %d, want 1", got.Attempts) + } + if got.LastError == nil || *got.LastError != "worker died" { + t.Fatalf("last_error: %v, want 'worker died'", got.LastError) + } +} + +func TestSweepOnce_DoesNotRequeueJobWithLiveWorker(t *testing.T) { + s, q, sw := setup(t) + ctx := context.Background() + + if err := q.Heartbeat(ctx, "w1", 30*time.Second); err != nil { + t.Fatal(err) + } + job, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, job.ID, "w1") + + if err := sw.SweepOnce(ctx); err != nil { + t.Fatalf("sweep: %v", err) + } + + got, _ := s.GetJob(ctx, job.ID) + if got.Status != store.StatusRunning { + t.Fatalf("status: %s, want running (untouched)", got.Status) + } +} + +func TestSweepOnce_PromotesDelayedRetryToRedisQueue(t *testing.T) { + s, q, sw := setup(t) + ctx := context.Background() + + job, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, job.ID, "w1") + _ = s.MarkFailed(ctx, job.ID, "boom", time.Now().Add(-1*time.Second)) + + depthBefore, _ := q.Depth(ctx, "test") + if err := sw.SweepOnce(ctx); err != nil { + t.Fatalf("sweep: %v", err) + } + depthAfter, _ := q.Depth(ctx, "test") + + if depthAfter != depthBefore+1 { + t.Fatalf("queue depth: before=%d after=%d, want +1", depthBefore, depthAfter) + } + + popped, err := q.BlockingPop(ctx, "test", 1*time.Second) + if err != nil { + t.Fatalf("pop: %v", err) + } + if popped != job.ID.String() { + t.Fatalf("popped %q, want %q", popped, job.ID.String()) + } +} + +func TestSweepOnce_LeavesFutureRetryInPostgresOnly(t *testing.T) { + s, q, sw := setup(t) + ctx := context.Background() + + job, _ := s.EnqueueJob(ctx, store.NewJob{Stage: "test"}) + _ = s.ClaimJob(ctx, job.ID, "w1") + _ = s.MarkFailed(ctx, job.ID, "boom", time.Now().Add(1*time.Hour)) + + depthBefore, _ := q.Depth(ctx, "test") + if err := sw.SweepOnce(ctx); err != nil { + t.Fatalf("sweep: %v", err) + } + depthAfter, _ := q.Depth(ctx, "test") + + if depthAfter != depthBefore { + t.Fatalf("queue depth changed: before=%d after=%d, want unchanged (future retry)", depthBefore, depthAfter) + } +}