diff --git a/runner/ratelimit_test.go b/runner/ratelimit_test.go new file mode 100644 index 00000000..c16cd5f3 --- /dev/null +++ b/runner/ratelimit_test.go @@ -0,0 +1,43 @@ +package runner + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestRunnerRateLimiter(t *testing.T) { + for _, tc := range []struct { + name string + options Options + }{ + {name: "per second", options: Options{RateLimit: 1}}, + {name: "per minute", options: Options{RateLimitMinute: 1}}, + } { + t.Run(tc.name, func(t *testing.T) { + r, err := New(&tc.options) + require.NoError(t, err) + t.Cleanup(r.Close) + require.True(t, r.ratelimiter.CanTake()) + r.ratelimiter.Take() + // The runner must observe the worker's token count, rather than a + // copy of the initial count made when constructing the runner. + require.Eventually(t, func() bool { + return !r.ratelimiter.CanTake() + }, 500*time.Millisecond, time.Millisecond) + }) + } +} + +func TestRunnerUnlimitedRateLimiter(t *testing.T) { + // Unlimited limiters replenish every millisecond. Repeated construction + // exercises initialization concurrently with replenishment under -race. + for range 10 { + r, err := New(&Options{}) + require.NoError(t, err) + r.ratelimiter.Take() + require.True(t, r.ratelimiter.CanTake()) + r.Close() + } +} diff --git a/runner/runner.go b/runner/runner.go index ea0cac3f..79dcd1a5 100644 --- a/runner/runner.go +++ b/runner/runner.go @@ -90,7 +90,7 @@ type Runner struct { hm *hybrid.HybridMap excludeCdn bool stats clistats.StatisticsClient - ratelimiter ratelimit.Limiter + ratelimiter *ratelimit.Limiter HostErrorsCache gcache.Cache[string, int] browser *Browser ditClassifier *dit.Classifier @@ -416,11 +416,11 @@ func New(options *Options) (*Runner, error) { runner.hm = hm if options.RateLimitMinute > 0 { - runner.ratelimiter = *ratelimit.New(context.Background(), uint(options.RateLimitMinute), time.Minute) + runner.ratelimiter = ratelimit.New(context.Background(), uint(options.RateLimitMinute), time.Minute) } else if options.RateLimit > 0 { - runner.ratelimiter = *ratelimit.New(context.Background(), uint(options.RateLimit), time.Second) + runner.ratelimiter = ratelimit.New(context.Background(), uint(options.RateLimit), time.Second) } else { - runner.ratelimiter = *ratelimit.NewUnlimited(context.Background()) + runner.ratelimiter = ratelimit.NewUnlimited(context.Background()) } if options.HostMaxErrors >= 0 {