package processing import ( "context" "errors" "sync" "sync/atomic" "testing" "time" "github.com/google/uuid" "github.com/jackc/pgx/v5" ) func TestClampProcessingWorkers(t *testing.T) { t.Parallel() cases := []struct { in, want int }{ {0, 1}, {-3, 1}, {1, 1}, {2, 2}, {MaxProcessingWorkers, MaxProcessingWorkers}, {MaxProcessingWorkers + 5, MaxProcessingWorkers}, } for _, tc := range cases { if got := ClampProcessingWorkers(tc.in); got != tc.want { t.Fatalf("ClampProcessingWorkers(%d)=%d want %d", tc.in, got, tc.want) } } } func TestJobSlotsBoundsConcurrentProcess(t *testing.T) { t.Parallel() const workers = 2 slots := NewJobSlots(workers) var inflight atomic.Int32 var maxInflight atomic.Int32 var started atomic.Int32 block := make(chan struct{}) claim := func(context.Context) (uuid.UUID, error) { return uuid.New(), nil } process := func(context.Context, uuid.UUID) error { n := inflight.Add(1) for { cur := maxInflight.Load() if n <= cur || maxInflight.CompareAndSwap(cur, n) { break } } defer inflight.Add(-1) <-block return nil } ctx := context.Background() n, err := slots.Fill(ctx, claim, process, nil) if err != nil { t.Fatalf("Fill: %v", err) } if n != workers { t.Fatalf("started=%d want %d", n, workers) } started.Store(int32(n)) // Extra TryStart must not exceed the bound while slots are busy. ok, err := slots.TryStart(ctx, claim, process, nil) if err != nil { t.Fatalf("TryStart while busy: %v", err) } if ok { t.Fatal("TryStart while busy: expected started=false") } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { if maxInflight.Load() == int32(workers) { break } time.Sleep(5 * time.Millisecond) } if got := maxInflight.Load(); got != int32(workers) { t.Fatalf("maxInflight=%d want %d", got, workers) } close(block) slots.Wait() if got := started.Load(); got != int32(workers) { t.Fatalf("started total=%d want %d", got, workers) } } func TestJobSlotsFillStopsOnNoRows(t *testing.T) { t.Parallel() slots := NewJobSlots(4) var claims atomic.Int32 claim := func(context.Context) (uuid.UUID, error) { if claims.Add(1) > 1 { return uuid.Nil, pgx.ErrNoRows } return uuid.New(), nil } process := func(context.Context, uuid.UUID) error { return nil } n, err := slots.Fill(context.Background(), claim, process, nil) if !errors.Is(err, pgx.ErrNoRows) { t.Fatalf("err=%v want ErrNoRows", err) } if n != 1 { t.Fatalf("started=%d want 1", n) } slots.Wait() } func TestJobSlotsOnDoneSeesProcessError(t *testing.T) { t.Parallel() slots := NewJobSlots(1) want := errors.New("boom") var gotErr error var wg sync.WaitGroup wg.Add(1) _, err := slots.TryStart( context.Background(), func(context.Context) (uuid.UUID, error) { return uuid.New(), nil }, func(context.Context, uuid.UUID) error { return want }, func(_ uuid.UUID, err error) { gotErr = err wg.Done() }, ) if err != nil { t.Fatalf("TryStart: %v", err) } wg.Wait() slots.Wait() if !errors.Is(gotErr, want) { t.Fatalf("onDone err=%v want %v", gotErr, want) } }