diff --git a/limiter.go b/limiter.go index 318af99..80c325c 100644 --- a/limiter.go +++ b/limiter.go @@ -75,7 +75,6 @@ type RateLimiter struct { // it increments the request count and returns false. This method does not send an HTTP response, // so the caller must handle the response themselves or use the RespondOnLimit() method instead. func (l *RateLimiter) OnLimit(w http.ResponseWriter, r *http.Request, key string) bool { - currentWindow := l.currentWindow(time.Now().UTC()) ctx := r.Context() limit := l.requestLimit @@ -83,10 +82,14 @@ func (l *RateLimiter) OnLimit(w http.ResponseWriter, r *http.Request, key string limit = val } setHeader(w, l.headers.Limit, strconv.Itoa(limit)) - setHeader(w, l.headers.Reset, strconv.FormatInt(currentWindow.Add(l.windowLength).Unix(), 10)) l.mu.Lock() - _, rateFloat, err := l.calculateRate(key, limit) + // A request may wait for the mutex across a window boundary. Sample the + // clock after acquiring it and use the same window for reading and writing. + now := time.Now().UTC() + currentWindow := l.currentWindow(now) + setHeader(w, l.headers.Reset, strconv.FormatInt(currentWindow.Add(l.windowLength).Unix(), 10)) + _, rateFloat, err := l.calculateRateAt(key, limit, now) if err != nil { l.mu.Unlock() l.onError(w, r, err) @@ -163,7 +166,10 @@ func (l *RateLimiter) currentWindow(t time.Time) time.Time { } func (l *RateLimiter) calculateRate(key string, requestLimit int) (bool, float64, error) { - now := time.Now().UTC() + return l.calculateRateAt(key, requestLimit, time.Now().UTC()) +} + +func (l *RateLimiter) calculateRateAt(key string, requestLimit int, now time.Time) (bool, float64, error) { currentWindow := l.currentWindow(now) previousWindow := currentWindow.Add(-l.windowLength) diff --git a/limiter_window_test.go b/limiter_window_test.go new file mode 100644 index 0000000..73ecfc3 --- /dev/null +++ b/limiter_window_test.go @@ -0,0 +1,120 @@ +package httprate + +import ( + "net/http" + "net/http/httptest" + "strconv" + "sync" + "testing" + "time" +) + +func TestOnLimitUsesSameWindowAfterWaitingForLock(t *testing.T) { + for _, customCounter := range []bool{false, true} { + name := "default counter" + if customCounter { + name = "custom counter" + } + t.Run(name, func(t *testing.T) { + const windowLength = 200 * time.Millisecond + var options []Option + if customCounter { + options = append(options, WithLimitCounter(NewLocalLimitCounter(windowLength))) + } + l := NewRateLimiter(10, windowLength, options...) + counter := &windowRecordingCounter{LimitCounter: l.Counter()} + l.limitCounter = counter + + // Hold the limiter lock as another request would while accessing a slow + // counter. The queued request reaches its first header before the lock. + l.mu.Lock() + locked := true + defer func() { + if locked { + l.mu.Unlock() + } + }() + entered := make(chan time.Time, 1) + w := &windowNotifyingWriter{ + ResponseRecorder: httptest.NewRecorder(), + onHeader: func() { entered <- l.currentWindow(time.Now().UTC()) }, + } + done := make(chan bool, 1) + go func() { done <- l.OnLimit(w, httptest.NewRequest("GET", "/", nil), "queued") }() + + var queuedWindow time.Time + select { + case queuedWindow = <-entered: + case <-time.After(5 * time.Second): + t.Fatal("request did not reach the limiter") + } + time.Sleep(time.Until(queuedWindow.Add(windowLength)) + 20*time.Millisecond) + // Simulate the request holding the lock recording another key in the + // new window before releasing the lock to the queued request. + currentWindow := l.currentWindow(time.Now().UTC()) + if err := counter.LimitCounter.Increment("other", currentWindow); err != nil { + t.Fatal(err) + } + l.mu.Unlock() + locked = false + select { + case limited := <-done: + if limited { + t.Fatal("first request for queued key was limited") + } + case <-time.After(5 * time.Second): + t.Fatal("request did not finish after releasing the lock") + } + + if !counter.getWindow.After(queuedWindow) { + t.Errorf("Get used stale window %v, queued in %v", counter.getWindow, queuedWindow) + } + if !counter.getWindow.Equal(counter.incrementWindow) { + t.Errorf("Get window %v differs from IncrementBy window %v", counter.getWindow, counter.incrementWindow) + } + if got, want := w.Header().Get("X-RateLimit-Reset"), strconv.FormatInt(counter.getWindow.Add(windowLength).Unix(), 10); got != want { + t.Errorf("reset header = %s, want %s", got, want) + } + curr, prev, err := counter.LimitCounter.Get("other", counter.getWindow, counter.getWindow.Add(-windowLength)) + // A heavily loaded scheduler can advance another window before the + // queued request resumes; allow the normal shift/expiry in that case. + wantCurr, wantPrev := 0, 0 + switch counter.getWindow.Sub(currentWindow) { + case 0: + wantCurr = 1 + case windowLength: + wantPrev = 1 + } + if err != nil || curr != wantCurr || prev != wantPrev { + t.Errorf("other key's counts = (%d, %d, %v), want (%d, %d, nil)", curr, prev, err, wantCurr, wantPrev) + } + }) + } +} + +type windowRecordingCounter struct { + LimitCounter + getWindow time.Time + incrementWindow time.Time +} + +func (c *windowRecordingCounter) Get(key string, currentWindow, previousWindow time.Time) (int, int, error) { + c.getWindow = currentWindow + return c.LimitCounter.Get(key, currentWindow, previousWindow) +} + +func (c *windowRecordingCounter) IncrementBy(key string, currentWindow time.Time, amount int) error { + c.incrementWindow = currentWindow + return c.LimitCounter.IncrementBy(key, currentWindow, amount) +} + +type windowNotifyingWriter struct { + *httptest.ResponseRecorder + onHeader func() + once sync.Once +} + +func (w *windowNotifyingWriter) Header() http.Header { + w.once.Do(w.onHeader) + return w.ResponseRecorder.Header() +}