diff --git a/internal/jobcompleter/job_completer_test.go b/internal/jobcompleter/job_completer_test.go index 749d57d39..de2bb3c2a 100644 --- a/internal/jobcompleter/job_completer_test.go +++ b/internal/jobcompleter/job_completer_test.go @@ -7,6 +7,7 @@ import ( "sync" "sync/atomic" "testing" + "testing/synctest" "time" "github.com/jackc/puddle/v2" @@ -177,111 +178,108 @@ func TestInlineJobCompleter_Subscribe(t *testing.T) { func TestInlineJobCompleter_Wait(t *testing.T) { t.Parallel() - testCompleterWait(t, func(schema string, exec riverdriver.Executor, subscribeChan SubscribeChan) JobCompleter { + testCompleterWait(t, func(t *testing.T, exec riverdriver.Executor, subscribeChan SubscribeChan) JobCompleter { + t.Helper() + return NewInlineCompleter(riversharedtest.BaseServiceArchetype(t), "", exec, &riverpilot.StandardPilot{}, subscribeChan) }) } -// TODO: Can we get rid of this test? It's pretty slow and it's not clear that -// it's testing anything particularly useful compared to the more thorough -// completer tests below. func TestAsyncJobCompleter_Complete(t *testing.T) { t.Parallel() - ctx := context.Background() + synctest.Test(t, func(t *testing.T) { + ctx := context.Background() - type jobInput struct { - // TODO: Try to get rid of containing the context in struct. It'd be - // better to pass it forward instead. - ctx context.Context //nolint:containedctx - jobID int64 - } - inputCh := make(chan jobInput) - resultCh := make(chan error) + type jobInput struct { + // TODO: Try to get rid of containing the context in struct. It'd be + // better to pass it forward instead. + ctx context.Context //nolint:containedctx + jobID int64 + } + inputCh := make(chan jobInput) + resultCh := make(chan error) - expectedErr := errors.New("an error from the completer") + expectedErr := errors.New("an error from the completer") - go func() { - riversharedtest.WaitOrTimeout(t, inputCh) - resultCh <- expectedErr - }() + go func() { + riversharedtest.WaitOrTimeout(t, inputCh) + resultCh <- expectedErr + }() - var ( - dbPool = riversharedtest.DBPool(ctx, t) - driver = riverpgxv5.New(dbPool) - schema = riverdbtest.TestSchema(ctx, t, driver, nil) - execMock = NewPartialExecutorMock(driver.GetExecutor()) - ) + execMock := &partialExecutorMock{} - execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { - require.Len(t, params.ID, 1) - inputCh <- jobInput{ctx: ctx, jobID: params.ID[0]} - err := <-resultCh - if err != nil { - return nil, err + execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { + require.Len(t, params.ID, 1) + inputCh <- jobInput{ctx: ctx, jobID: params.ID[0]} + err := <-resultCh + if err != nil { + return nil, err + } + return []*rivertype.JobRow{{ID: params.ID[0], State: params.State[0]}}, nil } - return []*rivertype.JobRow{{ID: params.ID[0], State: params.State[0]}}, nil - } - subscribeChan := make(chan []CompleterJobUpdated, 10) - completer := newAsyncCompleterWithConcurrency(riversharedtest.BaseServiceArchetype(t), schema, execMock, &riverpilot.StandardPilot{}, 2, subscribeChan) - completer.disableSleep = true - require.NoError(t, completer.Start(ctx)) - t.Cleanup(completer.Stop) + subscribeChan := make(chan []CompleterJobUpdated, 10) + completer := newAsyncCompleterWithConcurrency(riversharedtest.BaseServiceArchetype(t), "", execMock, &riverpilot.StandardPilot{}, 2, subscribeChan) + completer.disableSleep = true + require.NoError(t, completer.Start(ctx)) + t.Cleanup(completer.Stop) - // launch 4 completions, only 2 can be inline due to the concurrency limit: - for i := range int64(2) { - if err := completer.JobSetStateIfRunning(ctx, &jobstats.JobStatistics{}, riverdriver.JobSetStateCompleted(i, time.Now(), nil)); err != nil { - t.Errorf("expected nil err, got %v", err) - } - } - bgCompletionsStarted := make(chan struct{}) - go func() { - for i := int64(2); i < 4; i++ { + // launch 4 completions, only 2 can be inline due to the concurrency limit: + for i := range int64(2) { if err := completer.JobSetStateIfRunning(ctx, &jobstats.JobStatistics{}, riverdriver.JobSetStateCompleted(i, time.Now(), nil)); err != nil { t.Errorf("expected nil err, got %v", err) } } - close(bgCompletionsStarted) - }() + bgCompletionsStarted := make(chan struct{}) + go func() { + for i := int64(2); i < 4; i++ { + if err := completer.JobSetStateIfRunning(ctx, &jobstats.JobStatistics{}, riverdriver.JobSetStateCompleted(i, time.Now(), nil)); err != nil { + t.Errorf("expected nil err, got %v", err) + } + } + close(bgCompletionsStarted) + }() - expectCompletionInFlight := func() { - select { - case input := <-inputCh: - t.Logf("completion for %d in-flight", input.jobID) - case <-time.After(time.Second): - t.Fatalf("expected a completion to be in-flight") + expectCompletionInFlight := func() { + select { + case input := <-inputCh: + t.Logf("completion for %d in-flight", input.jobID) + case <-time.After(time.Second): + t.Fatalf("expected a completion to be in-flight") + } } - } - expectNoCompletionInFlight := func() { - select { - case input := <-inputCh: - t.Fatalf("unexpected completion for %d in-flight", input.jobID) - case <-time.After(500 * time.Millisecond): + expectNoCompletionInFlight := func() { + synctest.Wait() + select { + case input := <-inputCh: + t.Fatalf("unexpected completion for %d in-flight", input.jobID) + default: + } } - } - // two completions should be in-flight: - expectCompletionInFlight() - expectCompletionInFlight() + // two completions should be in-flight: + expectCompletionInFlight() + expectCompletionInFlight() - // A 3rd one shouldn't be in-flight due to the concurrency limit: - expectNoCompletionInFlight() + // A 3rd one shouldn't be in-flight due to the concurrency limit: + expectNoCompletionInFlight() - // Finish the first two completions: - resultCh <- nil - resultCh <- nil + // Finish the first two completions: + resultCh <- nil + resultCh <- nil - // The final two completions should now be in-flight: - <-bgCompletionsStarted - expectCompletionInFlight() - expectCompletionInFlight() + // The final two completions should now be in-flight: + <-bgCompletionsStarted + expectCompletionInFlight() + expectCompletionInFlight() - // A 5th one shouldn't be in-flight because we only started 4: - expectNoCompletionInFlight() + // A 5th one shouldn't be in-flight because we only started 4: + expectNoCompletionInFlight() - // Finish the final two completions: - resultCh <- nil - resultCh <- nil + // Finish the final two completions: + resultCh <- nil + resultCh <- nil + }) } func TestAsyncJobCompleter_CompleteDeletedJob(t *testing.T) { @@ -318,8 +316,10 @@ func TestAsyncJobCompleter_Subscribe(t *testing.T) { func TestAsyncJobCompleter_Wait(t *testing.T) { t.Parallel() - testCompleterWait(t, func(schema string, exec riverdriver.Executor, subscribeChan SubscribeChan) JobCompleter { - return newAsyncCompleterWithConcurrency(riversharedtest.BaseServiceArchetype(t), schema, exec, &riverpilot.StandardPilot{}, 4, subscribeChan) + testCompleterWait(t, func(t *testing.T, exec riverdriver.Executor, subscribeChan SubscribeChan) JobCompleter { + t.Helper() + + return newAsyncCompleterWithConcurrency(riversharedtest.BaseServiceArchetype(t), "", exec, &riverpilot.StandardPilot{}, 4, subscribeChan) }) } @@ -370,75 +370,75 @@ func testCompleterSubscribe(t *testing.T, constructor func(schema string, exec r } } -func testCompleterWait(t *testing.T, constructor func(schema string, exec riverdriver.Executor, subscribeChan SubscribeChan) JobCompleter) { +func testCompleterWait(t *testing.T, constructor func(t *testing.T, exec riverdriver.Executor, subscribeChan SubscribeChan) JobCompleter) { t.Helper() - ctx := context.Background() + synctest.Test(t, func(t *testing.T) { + ctx := context.Background() - var ( - dbPool = riversharedtest.DBPool(ctx, t) - driver = riverpgxv5.New(dbPool) - schema = riverdbtest.TestSchema(ctx, t, driver, nil) - execMock = NewPartialExecutorMock(driver.GetExecutor()) - ) + execMock := &partialExecutorMock{} - resultCh := make(chan struct{}) - completeStartedCh := make(chan struct{}) - execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { - completeStartedCh <- struct{}{} - <-resultCh - results := make([]*rivertype.JobRow, len(params.ID)) - for i := range params.ID { - results[i] = &rivertype.JobRow{ID: params.ID[i], State: rivertype.JobStateCompleted} + resultCh := make(chan struct{}) + completeStartedCh := make(chan struct{}) + execMock.JobSetStateIfRunningManyFunc = func(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { + completeStartedCh <- struct{}{} + <-resultCh + results := make([]*rivertype.JobRow, len(params.ID)) + for i := range params.ID { + results[i] = &rivertype.JobRow{ID: params.ID[i], State: rivertype.JobStateCompleted} + } + return results, nil } - return results, nil - } - subscribeCh := make(chan []CompleterJobUpdated, 100) + subscribeCh := make(chan []CompleterJobUpdated, 100) - completer := constructor(schema, execMock, subscribeCh) - require.NoError(t, completer.Start(ctx)) + completer := constructor(t, execMock, subscribeCh) + require.NoError(t, completer.Start(ctx)) - // launch 4 completions: - for i := range 4 { - go func() { - require.NoError(t, completer.JobSetStateIfRunning(ctx, &jobstats.JobStatistics{}, riverdriver.JobSetStateCompleted(int64(i), time.Now(), nil))) - }() - <-completeStartedCh // wait for func to actually start - } + // launch 4 completions: + for i := range 4 { + go func() { + require.NoError(t, completer.JobSetStateIfRunning(ctx, &jobstats.JobStatistics{}, riverdriver.JobSetStateCompleted(int64(i), time.Now(), nil))) + }() + <-completeStartedCh // wait for func to actually start + } - // Give one completion a signal to finish, there should be 3 remaining in-flight: - resultCh <- struct{}{} + // Give one completion a signal to finish, there should be 3 remaining in-flight: + resultCh <- struct{}{} - waitDone := make(chan struct{}) - go func() { - completer.Stop() - close(waitDone) - }() + waitDone := make(chan struct{}) + go func() { + completer.Stop() + close(waitDone) + }() - select { - case <-waitDone: - t.Fatalf("expected Wait to block until all jobs are complete, but it returned when there should be three remaining") - case <-time.After(100 * time.Millisecond): - } + synctest.Wait() + select { + case <-waitDone: + t.Fatalf("expected Wait to block until all jobs are complete, but it returned when there should be three remaining") + default: + } - // Get us down to one in-flight completion: - resultCh <- struct{}{} - resultCh <- struct{}{} + // Get us down to one in-flight completion: + resultCh <- struct{}{} + resultCh <- struct{}{} - select { - case <-waitDone: - t.Fatalf("expected Wait to block until all jobs are complete, but it returned when there should be one remaining") - case <-time.After(100 * time.Millisecond): - } + synctest.Wait() + select { + case <-waitDone: + t.Fatalf("expected Wait to block until all jobs are complete, but it returned when there should be one remaining") + default: + } - // Finish the last one: - resultCh <- struct{}{} + // Finish the last one: + resultCh <- struct{}{} - select { - case <-waitDone: - case <-time.After(100 * time.Millisecond): - t.Errorf("expected Wait to return after all jobs are complete") - } + synctest.Wait() + select { + case <-waitDone: + default: + t.Errorf("expected Wait to return after all jobs are complete") + } + }) } func TestAsyncCompleter(t *testing.T) { diff --git a/internal/maintenance/periodic_job_enqueuer_test.go b/internal/maintenance/periodic_job_enqueuer_test.go index ae99f64f9..28c1be8e3 100644 --- a/internal/maintenance/periodic_job_enqueuer_test.go +++ b/internal/maintenance/periodic_job_enqueuer_test.go @@ -9,6 +9,7 @@ import ( "strings" "sync" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/require" @@ -696,46 +697,49 @@ func TestPeriodicJobEnqueuer(t *testing.T) { require.Len(t, svc.periodicJobs, 1) }) - // To suss out any race conditions in the add/remove/clear/run loop code, - // and interactions between them. + // Exercise concurrent changes to the in-memory periodic job registry. t.Run("AddRemoveStress", func(t *testing.T) { t.Parallel() - svc, _ := setup(t) + synctest.Test(t, func(t *testing.T) { + // This test only mutates the in-memory registry; no service is started. + svc, err := NewPeriodicJobEnqueuer(riversharedtest.BaseServiceArchetype(t), &PeriodicJobEnqueuerConfig{}, nil) + require.NoError(t, err) - var wg sync.WaitGroup + var wg sync.WaitGroup - randomSleep := func() { - time.Sleep(time.Duration(randutil.IntBetween(1, 5)) * time.Millisecond) - } + randomSleepFunc := func() { + time.Sleep(time.Duration(randutil.IntBetween(1, 5)) * time.Millisecond) + } - for i := range 10 { - wg.Add(1) + for i := range 10 { + wg.Add(1) - jobBaseName := fmt.Sprintf("periodic_job_1ms_%02d", i) + jobBaseName := fmt.Sprintf("periodic_job_1ms_%02d", i) - go func() { - defer wg.Done() + go func() { + defer wg.Done() - for range 50 { - handle, err := svc.AddSafely(&PeriodicJob{ScheduleFunc: periodicIntervalSchedule(time.Millisecond), ConstructorFunc: jobConstructorFunc(jobBaseName, false)}) - require.NoError(t, err) - randomSleep() + for range 50 { + handle, err := svc.AddSafely(&PeriodicJob{ScheduleFunc: periodicIntervalSchedule(time.Millisecond), ConstructorFunc: jobConstructorFunc(jobBaseName, false)}) + require.NoError(t, err) + randomSleepFunc() - _, err = svc.AddSafely(&PeriodicJob{ScheduleFunc: periodicIntervalSchedule(time.Millisecond), ConstructorFunc: jobConstructorFunc(jobBaseName+"_second", false)}) - require.NoError(t, err) - randomSleep() + _, err = svc.AddSafely(&PeriodicJob{ScheduleFunc: periodicIntervalSchedule(time.Millisecond), ConstructorFunc: jobConstructorFunc(jobBaseName+"_second", false)}) + require.NoError(t, err) + randomSleepFunc() - svc.Remove(handle) - randomSleep() + svc.Remove(handle) + randomSleepFunc() - svc.Clear() - randomSleep() - } - }() - } + svc.Clear() + randomSleepFunc() + } + }() + } - wg.Wait() + wg.Wait() + }) }) t.Run("NoJobsConfigured", func(t *testing.T) { diff --git a/internal/notifylimiter/limiter_test.go b/internal/notifylimiter/limiter_test.go index f69ceafed..ca4eb2147 100644 --- a/internal/notifylimiter/limiter_test.go +++ b/internal/notifylimiter/limiter_test.go @@ -1,8 +1,10 @@ package notifylimiter import ( + "sync" "sync/atomic" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/require" @@ -13,21 +15,16 @@ import ( func TestLimiter(t *testing.T) { t.Parallel() - type testBundle struct{} + setup := func(t *testing.T) *Limiter { + t.Helper() - setup := func() (*Limiter, *testBundle) { - bundle := &testBundle{} - - archetype := riversharedtest.BaseServiceArchetype(t) - limiter := NewLimiter(archetype, 10*time.Millisecond) - - return limiter, bundle + return NewLimiter(riversharedtest.BaseServiceArchetype(t), 10*time.Millisecond) } t.Run("OnlySendsOncePerWaitDuration", func(t *testing.T) { t.Parallel() - limiter, _ := setup() + limiter := setup(t) now := time.Now() limiter.Time.StubNow(now) @@ -60,49 +57,42 @@ func TestLimiter(t *testing.T) { t.Run("ConcurrentAccessStressTest", func(t *testing.T) { t.Parallel() - doneCh := make(chan struct{}) - t.Cleanup(func() { close(doneCh) }) + synctest.Test(t, func(t *testing.T) { + limiter := setup(t) - limiter, _ := setup() - now := time.Now() - limiter.Time.StubNow(now) - - counters := make(map[string]*atomic.Int64) - for _, topic := range []string{"a", "b", "c"} { - counters[topic] = &atomic.Int64{} - } + counters := make(map[string]*atomic.Int64) + for _, topic := range []string{"a", "b", "c"} { + counters[topic] = &atomic.Int64{} + } - signalContinuously := func(topic string) { - for { - select { - case <-doneCh: - return - default: - shouldTrigger := limiter.ShouldTrigger(topic) - if shouldTrigger { - counters[topic].Add(1) + // Bounded rounds exercise concurrent access without busy loops that + // would prevent the synctest clock from advancing. + signalConcurrentlyFunc := func() { + var wg sync.WaitGroup + for topic := range counters { + for range 10 { + wg.Go(func() { + for range 100 { + if limiter.ShouldTrigger(topic) { + counters[topic].Add(1) + } + } + }) } } + wg.Wait() } - } - go signalContinuously("a") - go signalContinuously("b") - go signalContinuously("c") - - // Duration doesn't really matter here, just need time for these all to fire - // a bit: - <-time.After(100 * time.Millisecond) - require.Equal(t, int64(1), counters["a"].Load()) - require.Equal(t, int64(1), counters["b"].Load()) - require.Equal(t, int64(1), counters["c"].Load()) - - limiter.Time.StubNow(now.Add(11 * time.Millisecond)) - - <-time.After(100 * time.Millisecond) + signalConcurrentlyFunc() + for _, counter := range counters { + require.Equal(t, int64(1), counter.Load()) + } - require.Equal(t, int64(2), counters["a"].Load()) - require.Equal(t, int64(2), counters["b"].Load()) - require.Equal(t, int64(2), counters["c"].Load()) + time.Sleep(11 * time.Millisecond) + signalConcurrentlyFunc() + for _, counter := range counters { + require.Equal(t, int64(2), counter.Load()) + } + }) }) } diff --git a/internal/util/chanutil/debounced_chan_test.go b/internal/util/chanutil/debounced_chan_test.go index 10df45891..ac36b6bf5 100644 --- a/internal/util/chanutil/debounced_chan_test.go +++ b/internal/util/chanutil/debounced_chan_test.go @@ -12,97 +12,90 @@ import ( func TestDebouncedChan_TriggersImmediately(t *testing.T) { t.Parallel() - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - debouncedChan := NewDebouncedChan(ctx, 200*time.Millisecond, true) - go debouncedChan.Call() - - select { - case <-debouncedChan.C(): - case <-time.After(50 * time.Millisecond): - t.Fatal("timed out waiting for debounced chan to trigger") - } - - // shouldn't trigger immediately again - go debouncedChan.Call() - select { - case <-debouncedChan.C(): - t.Fatal("received from debounced chan unexpectedly") - case <-time.After(50 * time.Millisecond): - } - - var wg sync.WaitGroup - wg.Add(5) - for range 5 { - go func() { - debouncedChan.Call() - wg.Done() - }() - } - wg.Wait() - - // should trigger again after debounce period - select { - case <-debouncedChan.C(): - case <-time.After(250 * time.Millisecond): - t.Fatal("timed out waiting for debounced chan to trigger") - } - - // shouldn't trigger immediately again - select { - case <-debouncedChan.C(): - t.Fatal("received from debounced chan unexpectedly") - case <-time.After(50 * time.Millisecond): - } + + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + const cooldown = 200 * time.Millisecond + debouncedChan := NewDebouncedChan(ctx, cooldown, true) + go debouncedChan.Call() + synctest.Wait() + + require.Len(t, debouncedChan.C(), 1) + <-debouncedChan.C() + + // Concurrent calls during the cooldown coalesce into one trailing event. + var wg sync.WaitGroup + for range 5 { + wg.Go(debouncedChan.Call) + } + wg.Wait() + synctest.Wait() + require.Empty(t, debouncedChan.C()) + + time.Sleep(cooldown - time.Nanosecond) + synctest.Wait() + require.Empty(t, debouncedChan.C()) + + time.Sleep(time.Nanosecond) + synctest.Wait() + require.Len(t, debouncedChan.C(), 1) + <-debouncedChan.C() + + // No further calls means no additional trailing event. + time.Sleep(cooldown) + synctest.Wait() + require.Empty(t, debouncedChan.C()) + }) } func TestDebouncedChan_OnlyBuffersOneEvent(t *testing.T) { t.Parallel() - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - debouncedChan := NewDebouncedChan(ctx, 100*time.Millisecond, true) - debouncedChan.Call() - time.Sleep(150 * time.Millisecond) - debouncedChan.Call() - - select { - case <-debouncedChan.C(): - case <-time.After(20 * time.Millisecond): - t.Fatal("timed out waiting for debounced chan to trigger") - } - - // shouldn't trigger immediately again - select { - case <-debouncedChan.C(): - t.Fatal("received from debounced chan unexpectedly") - case <-time.After(20 * time.Millisecond): - } + + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + const cooldown = 100 * time.Millisecond + debouncedChan := NewDebouncedChan(ctx, cooldown, true) + debouncedChan.Call() + time.Sleep(cooldown) + synctest.Wait() + debouncedChan.Call() + synctest.Wait() + + require.Len(t, debouncedChan.C(), 1) + <-debouncedChan.C() + + time.Sleep(cooldown) + synctest.Wait() + require.Empty(t, debouncedChan.C()) + }) } func TestDebouncedChan_SendLeadingDisabled(t *testing.T) { t.Parallel() - ctx := context.Background() - - debouncedChan := NewDebouncedChan(ctx, 100*time.Millisecond, false) - debouncedChan.Call() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() - // Expect nothing right away because sendLeading is disabled. - select { - case <-debouncedChan.C(): - t.Fatal("received from debounced chan unexpectedly") - case <-time.After(20 * time.Millisecond): - } + const cooldown = 100 * time.Millisecond + debouncedChan := NewDebouncedChan(ctx, cooldown, false) + debouncedChan.Call() + synctest.Wait() + require.Empty(t, debouncedChan.C()) - time.Sleep(100 * time.Millisecond) + time.Sleep(cooldown - time.Nanosecond) + synctest.Wait() + require.Empty(t, debouncedChan.C()) - select { - case <-debouncedChan.C(): - case <-time.After(20 * time.Millisecond): - t.Fatal("timed out waiting for debounced chan to trigger") - } + time.Sleep(time.Nanosecond) + synctest.Wait() + require.Len(t, debouncedChan.C(), 1) + <-debouncedChan.C() + }) } func TestDebouncedChan_ContinuousOperation(t *testing.T) { diff --git a/rivershared/util/serviceutil/service_util_test.go b/rivershared/util/serviceutil/service_util_test.go index 8b991d354..0e75c9d19 100644 --- a/rivershared/util/serviceutil/service_util_test.go +++ b/rivershared/util/serviceutil/service_util_test.go @@ -3,6 +3,7 @@ package serviceutil import ( "context" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/require" @@ -14,28 +15,30 @@ func TestCancellableSleep(t *testing.T) { testCancellableSleep := func(t *testing.T, startSleepFunc func(ctx context.Context) <-chan struct{}) { t.Helper() - ctx := context.Background() - - ctx, cancel := context.WithCancel(ctx) - t.Cleanup(cancel) - - sleepDone := startSleepFunc(ctx) - - // Wait a very nominal amount of time just to make sure that some sleep is - // actually happening. - select { - case <-sleepDone: - require.FailNow(t, "Sleep returned sooner than expected") - case <-time.After(50 * time.Millisecond): - } - - cancel() - - select { - case <-sleepDone: - case <-time.After(50 * time.Millisecond): - require.FailNow(t, "Timed out waiting for sleep to finish after cancel") - } + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + sleepDone := startSleepFunc(ctx) + synctest.Wait() + + // Advance to just before the sleep would finish naturally. + time.Sleep(5*time.Second - time.Nanosecond) + synctest.Wait() + select { + case <-sleepDone: + t.Fatal("Sleep returned sooner than expected") + default: + } + + cancel() + synctest.Wait() + select { + case <-sleepDone: + default: + t.Fatal("Sleep did not finish after cancel") + } + }) } // Starts sleep for sleep functions that don't return a channel, returning a diff --git a/rivershared/util/timeutil/time_util_test.go b/rivershared/util/timeutil/time_util_test.go index 4ca1ba9a8..92a66cca5 100644 --- a/rivershared/util/timeutil/time_util_test.go +++ b/rivershared/util/timeutil/time_util_test.go @@ -3,11 +3,11 @@ package timeutil_test import ( "context" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/require" - "github.com/riverqueue/river/rivershared/riversharedtest" "github.com/riverqueue/river/rivershared/util/timeutil" ) @@ -20,28 +20,52 @@ func TestSecondsAsDuration(t *testing.T) { func TestTickerWithInitialTick(t *testing.T) { t.Parallel() - ctx := context.Background() - t.Run("TicksImmediately", func(t *testing.T) { t.Parallel() - ctx, cancel := context.WithCancel(ctx) - t.Cleanup(cancel) + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() - ticker := timeutil.NewTickerWithInitialTick(ctx, 1*time.Hour) - riversharedtest.WaitOrTimeout(t, ticker.C) + now := time.Now() + ticker := timeutil.NewTickerWithInitialTick(ctx, time.Hour) + synctest.Wait() + select { + case tick := <-ticker.C: + require.Equal(t, now, tick) + default: + t.Fatal("Initial tick was not immediate") + } + }) }) t.Run("TicksPeriodically", func(t *testing.T) { t.Parallel() - ctx, cancel := context.WithCancel(ctx) - t.Cleanup(cancel) + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + const interval = 100 * time.Microsecond + now := time.Now() + ticker := timeutil.NewTickerWithInitialTick(ctx, interval) + synctest.Wait() + require.Equal(t, now, <-ticker.C) - ticker := timeutil.NewTickerWithInitialTick(ctx, 100*time.Microsecond) - for i := range 10 { - t.Logf("Waiting on tick %d", i) - riversharedtest.WaitOrTimeout(t, ticker.C) - } + for range 9 { + time.Sleep(interval - time.Nanosecond) + synctest.Wait() + require.Empty(t, ticker.C) + time.Sleep(time.Nanosecond) + synctest.Wait() + now = now.Add(interval) + select { + case tick := <-ticker.C: + require.Equal(t, now, tick) + default: + t.Fatal("Periodic tick was not delivered") + } + } + }) }) }