Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 68 additions & 0 deletions cmd/sweeper/main.go
Original file line number Diff line number Diff line change
@@ -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
}
8 changes: 8 additions & 0 deletions internal/queue/queue.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
26 changes: 26 additions & 0 deletions internal/queue/queue_integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
43 changes: 43 additions & 0 deletions internal/store/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
61 changes: 61 additions & 0 deletions internal/store/store_integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
95 changes: 95 additions & 0 deletions internal/sweeper/sweeper.go
Original file line number Diff line number Diff line change
@@ -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
}
Loading
Loading