package wait import ( "context" "errors" "testing" "time" ) // fakeClock advances instantly on Sleep so the loop runs with no real delay. type fakeClock struct { t time.Time cancelAt time.Duration // if >0, cancel the run once elapsed reaches this cancel context.CancelFunc start time.Time sleepCall int } func newFakeClock() *fakeClock { start := time.Unix(0, 0) return &fakeClock{t: start, start: start} } func (c *fakeClock) Now() time.Time { return c.t } func (c *fakeClock) Sleep(ctx context.Context, d time.Duration) error { c.sleepCall++ c.t = c.t.Add(d) if c.cancelAt > 0 && c.t.Sub(c.start) >= c.cancelAt && c.cancel != nil { c.cancel() } return ctx.Err() } var errDown = errors.New("connection refused") // failNThenOK returns an AttemptFunc that fails the first n calls then succeeds. func failNThenOK(n int, calls *int) AttemptFunc { return func(ctx context.Context) error { *calls++ if *calls <= n { return errDown } return nil } } func noFail(int, error, time.Duration, time.Duration, time.Duration) {} func TestSucceedsFirstAttempt(t *testing.T) { calls := 0 res := Run(context.Background(), Params{Timeout: time.Minute, Interval: 2 * time.Second, ConnectTimeout: time.Second}, failNThenOK(0, &calls), noFail, newFakeClock()) if !res.OK || res.Attempts != 1 { t.Fatalf("want OK after 1 attempt, got %+v", res) } } func TestWaitsThenSucceeds(t *testing.T) { calls := 0 clk := newFakeClock() failures := 0 res := Run(context.Background(), Params{Timeout: time.Minute, Interval: 2 * time.Second, ConnectTimeout: time.Second}, failNThenOK(3, &calls), func(int, error, time.Duration, time.Duration, time.Duration) { failures++ }, clk) if !res.OK { t.Fatalf("want OK, got %+v", res) } if res.Attempts != 4 { t.Errorf("attempts = %d, want 4", res.Attempts) } if failures != 3 { t.Errorf("onFail called %d times, want 3", failures) } // 3 sleeps of 2s each. if got := res.Elapsed; got != 6*time.Second { t.Errorf("elapsed = %v, want 6s", got) } } func TestTimesOut(t *testing.T) { calls := 0 alwaysFail := func(ctx context.Context) error { calls++; return errDown } res := Run(context.Background(), Params{Timeout: 10 * time.Second, Interval: 3 * time.Second, ConnectTimeout: time.Second}, alwaysFail, noFail, newFakeClock()) if res.OK || !res.TimedOut { t.Fatalf("want timeout, got %+v", res) } if !errors.Is(res.LastErr, errDown) { t.Errorf("LastErr = %v, want errDown", res.LastErr) } // Deadline 10s, interval 3s: attempts at 0,3,6,9, then next check at ~12s > deadline. if res.Attempts < 3 { t.Errorf("attempts = %d, want several before timeout", res.Attempts) } } func TestWaitForeverEventuallySucceeds(t *testing.T) { calls := 0 res := Run(context.Background(), Params{Timeout: 0, Interval: time.Second, ConnectTimeout: time.Second}, failNThenOK(100, &calls), noFail, newFakeClock()) if !res.OK || res.Attempts != 101 { t.Fatalf("want OK after 101 attempts with no timeout, got %+v", res) } } func TestCancelledDuringSleep(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) clk := newFakeClock() clk.cancelAt = 4 * time.Second clk.cancel = cancel calls := 0 alwaysFail := func(ctx context.Context) error { calls++; return errDown } res := Run(ctx, Params{Timeout: time.Hour, Interval: 2 * time.Second, ConnectTimeout: time.Second}, alwaysFail, noFail, clk) if !res.Cancelled { t.Fatalf("want Cancelled, got %+v", res) } }