diff --git a/runner/runner.go b/runner/runner.go index ea0cac3f..4aa64db2 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 { @@ -917,7 +917,9 @@ func (r *Runner) Close() { // nolint:errcheck // ignore r.hm.Close() r.hp.Dialer.Close() - r.ratelimiter.Stop() + if r.ratelimiter != nil { + r.ratelimiter.Stop() + } if r.options.HostMaxErrors >= 0 { r.HostErrorsCache.Purge() diff --git a/runner/runner_test.go b/runner/runner_test.go index 3f9d5e55..940da70a 100644 --- a/runner/runner_test.go +++ b/runner/runner_test.go @@ -1166,3 +1166,32 @@ func TestHandshakeThenCloseKeepsPlainResult(t *testing.T) { require.EqualValues(t, 1, atomic.LoadInt64(&tlsHandshakes), "the failed HTTPS request must complete its TLS handshake first") } + +func TestRunner_RateLimiterInitialization(t *testing.T) { + t.Run("default unlimited rate limiter is non-nil pointer", func(t *testing.T) { + r, err := New(&Options{}) + require.NoError(t, err) + defer r.Close() + + require.NotNil(t, r.ratelimiter) + r.ratelimiter.Take() + }) + + t.Run("per-second rate limiter is non-nil pointer", func(t *testing.T) { + r, err := New(&Options{RateLimit: 50}) + require.NoError(t, err) + defer r.Close() + + require.NotNil(t, r.ratelimiter) + r.ratelimiter.Take() + }) + + t.Run("per-minute rate limiter is non-nil pointer", func(t *testing.T) { + r, err := New(&Options{RateLimitMinute: 60}) + require.NoError(t, err) + defer r.Close() + + require.NotNil(t, r.ratelimiter) + r.ratelimiter.Take() + }) +}