Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 10 additions & 4 deletions limiter.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,18 +75,21 @@ 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
if val := getRequestLimit(ctx); val > 0 {
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)
Expand Down Expand Up @@ -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)

Expand Down
120 changes: 120 additions & 0 deletions limiter_window_test.go
Original file line number Diff line number Diff line change
@@ -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()
}
Loading