diff --git a/surf/limiter_test.go b/surf/limiter_test.go index 3ad737a..0b55639 100644 --- a/surf/limiter_test.go +++ b/surf/limiter_test.go @@ -5,6 +5,8 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" + "sync/atomic" "testing" "time" @@ -14,6 +16,115 @@ import ( "git.golem15.com/golem15/summercms/party" ) +func TestMemoryStoreAtomicAttempt(t *testing.T) { + s := NewMemoryStore(0) + decay := 30 * time.Millisecond + + allowed, attempts, retryAfter := s.Attempt("k", 1, decay) + if !allowed || attempts != 1 || retryAfter <= 0 { + t.Fatalf("first attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter) + } + allowed, attempts, retryAfter = s.Attempt("k", 1, decay) + if allowed || attempts != 1 || retryAfter <= 0 { + t.Fatalf("denied attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter) + } + + time.Sleep(decay + 10*time.Millisecond) + allowed, attempts, retryAfter = s.Attempt("k", 1, decay) + if !allowed || attempts != 1 || retryAfter <= 0 { + t.Fatalf("fresh-window attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter) + } +} + +func TestMemoryStoreConcurrentAttempt(t *testing.T) { + const workers = 32 + s := NewMemoryStore(0) + ready := sync.WaitGroup{} + ready.Add(workers) + start := make(chan struct{}) + results := make(chan bool, workers) + + var workersDone sync.WaitGroup + workersDone.Add(workers) + for range workers { + go func() { + defer workersDone.Done() + ready.Done() + <-start + allowed, attempts, _ := s.Attempt("shared", 1, time.Minute) + if attempts != 1 { + t.Errorf("attempts = %d, want 1", attempts) + } + results <- allowed + }() + } + ready.Wait() + close(start) + workersDone.Wait() + close(results) + + allowed := 0 + for result := range results { + if result { + allowed++ + } + } + if allowed != 1 { + t.Fatalf("allowed = %d, want 1", allowed) + } +} + +func TestFixedWindowLimiterConcurrentMaxOne(t *testing.T) { + const workers = 32 + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + var handlerCalls atomic.Int32 + h := lim.Middleware("1,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + handlerCalls.Add(1) + w.WriteHeader(http.StatusNoContent) + })) + + ready := sync.WaitGroup{} + ready.Add(workers) + start := make(chan struct{}) + statuses := make(chan int, workers) + var workersDone sync.WaitGroup + workersDone.Add(workers) + for range workers { + go func() { + defer workersDone.Done() + req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + req.RemoteAddr = "192.0.2.1:1234" + ready.Done() + <-start + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code == http.StatusTooManyRequests && rec.Body.String() != tooManyAttemptsBody { + t.Errorf("429 body = %q", rec.Body.String()) + } + statuses <- rec.Code + }() + } + ready.Wait() + close(start) + workersDone.Wait() + close(statuses) + + successes, denied := 0, 0 + for status := range statuses { + switch status { + case http.StatusNoContent: + successes++ + case http.StatusTooManyRequests: + denied++ + default: + t.Errorf("unexpected status %d", status) + } + } + if successes != 1 || denied != workers-1 || handlerCalls.Load() != 1 { + t.Fatalf("successes=%d denied=%d handlerCalls=%d", successes, denied, handlerCalls.Load()) + } +} + func TestFixedWindowLimiterMemoryStoreWindow(t *testing.T) { s := NewMemoryStore(0) decay := 80 * time.Millisecond @@ -209,12 +320,12 @@ func TestFixedWindowLimiterInlineThrottleKeys(t *testing.T) { } }) - t.Run("anonymous same IP+Host", func(t *testing.T) { + t.Run("anonymous same IP different Host", func(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) h := lim.Middleware("1,1")(okHandler) - req1 := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + req1 := httptest.NewRequest(http.MethodGet, "http://first.example/x", nil) req1.RemoteAddr = "192.0.2.1:1" - req2 := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + req2 := httptest.NewRequest(http.MethodGet, "http://second.example/x", nil) req2.RemoteAddr = "192.0.2.1:9" first := httptest.NewRecorder() h.ServeHTTP(first, req1) @@ -224,7 +335,29 @@ func TestFixedWindowLimiterInlineThrottleKeys(t *testing.T) { second := httptest.NewRecorder() h.ServeHTTP(second, req2) if second.Code != http.StatusTooManyRequests { - t.Fatalf("anonymous same IP+Host should share key, status = %d", second.Code) + t.Fatalf("anonymous same IP with rotated Host should share key, status = %d", second.Code) + } + }) + + t.Run("anonymous inline policies share a domainless key", func(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + twoPerMinute := lim.Middleware("2,1")(okHandler) + onePerMinute := lim.Middleware("1,1")(okHandler) + request := func(h http.Handler) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + req.RemoteAddr = "192.0.2.9:1234" + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + return rec + } + if rec := request(twoPerMinute); rec.Code != http.StatusOK { + t.Fatalf("first throttle:2,1 status = %d", rec.Code) + } + if rec := request(twoPerMinute); rec.Code != http.StatusOK { + t.Fatalf("second throttle:2,1 status = %d", rec.Code) + } + if rec := request(onePerMinute); rec.Code != http.StatusTooManyRequests { + t.Fatalf("throttle:1,1 after shared exhaustion status = %d, want 429", rec.Code) } }) }