package surf import ( "encoding/json" "net/http" "net/http/httptest" "strings" "sync" "sync/atomic" "testing" "time" "git.golem15.com/golem15/summercms/modules/backpack" "git.golem15.com/golem15/summercms/modules/bouncer" "git.golem15.com/golem15/summercms/modules/pact" "git.golem15.com/golem15/summercms/modules/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 if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 1 { t.Fatalf("attempt 1 = allowed %v, attempts %d", allowed, attempts) } if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 2 { t.Fatalf("attempt 2 = allowed %v, attempts %d", allowed, attempts) } if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 3 { t.Fatalf("attempt 3 = allowed %v, attempts %d", allowed, attempts) } if allowed, attempts, _ := s.Attempt("k", 3, decay); allowed || attempts != 3 { t.Fatalf("denied attempt = allowed %v, attempts %d", allowed, attempts) } time.Sleep(decay + 20*time.Millisecond) if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 1 { t.Fatalf("fresh-window attempt = allowed %v, attempts %d", allowed, attempts) } } func TestFixedWindowLimiterMemoryStoreFirstHitWins(t *testing.T) { s := NewMemoryStore(0) decay := 200 * time.Millisecond if allowed, _, _ := s.Attempt("k", 2, decay); !allowed { t.Fatal("first attempt denied") } time.Sleep(120 * time.Millisecond) if allowed, _, _ := s.Attempt("k", 2, decay); !allowed { t.Fatal("second attempt denied") } time.Sleep(100 * time.Millisecond) if allowed, attempts, _ := s.Attempt("k", 1, decay); !allowed || attempts != 1 { t.Fatal("window was extended by a later hit") } } func TestFixedWindowLimiterSuccessHeaders(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) h := lim.Middleware("3,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(`{"ok":true}`)) })) rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) h.ServeHTTP(rec, req) if rec.Code != http.StatusOK { t.Fatalf("status = %d", rec.Code) } if rec.Header().Get("X-RateLimit-Limit") != "3" { t.Fatalf("limit = %q", rec.Header().Get("X-RateLimit-Limit")) } if rec.Header().Get("X-RateLimit-Remaining") != "2" { t.Fatalf("remaining = %q", rec.Header().Get("X-RateLimit-Remaining")) } if rec.Header().Get("Retry-After") != "" || rec.Header().Get("X-RateLimit-Reset") != "" { t.Fatal("Retry-After / X-RateLimit-Reset on success") } } func TestFixedWindowLimiterTooManyAttemptsHeaders(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) h := lim.Middleware("1,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })) req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) ok := httptest.NewRecorder() h.ServeHTTP(ok, req) if ok.Code != http.StatusOK { t.Fatalf("first status = %d", ok.Code) } rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusTooManyRequests { t.Fatalf("status = %d", rec.Code) } if rec.Header().Get("Retry-After") == "" { t.Fatal("missing Retry-After") } if rec.Header().Get("X-RateLimit-Reset") == "" { t.Fatal("missing X-RateLimit-Reset") } if rec.Header().Get("X-RateLimit-Limit") != "1" { t.Fatalf("limit = %q", rec.Header().Get("X-RateLimit-Limit")) } if rec.Header().Get("X-RateLimit-Remaining") != "0" { t.Fatalf("remaining = %q", rec.Header().Get("X-RateLimit-Remaining")) } if rec.Body.String() != `{"message":"Too Many Attempts."}` { t.Fatalf("body = %q", rec.Body.String()) } var payload map[string]any if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { t.Fatal(err) } if payload["message"] != "Too Many Attempts." { t.Fatalf("payload = %v", payload) } } func TestFixedWindowLimiterStackedBuckets(t *testing.T) { okHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) t.Run("exhaust A", func(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) if err := lim.RegisterBucket("demo", "bucket-a", Bucket{ Max: 1, Decay: time.Minute, Key: func(*http.Request) string { return "a" }, }); err != nil { t.Fatal(err) } if err := lim.RegisterBucket("demo", "bucket-b", Bucket{ Max: 100, Decay: time.Minute, Key: func(*http.Request) string { return "b" }, }); err != nil { t.Fatal(err) } h := stackThrottle(lim, okHandler, "bucket-a", "bucket-b") req := httptest.NewRequest(http.MethodGet, "/x", nil) first := httptest.NewRecorder() h.ServeHTTP(first, req) if first.Code != http.StatusOK { t.Fatalf("first = %d", first.Code) } rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusTooManyRequests { t.Fatalf("exhausted A status = %d", rec.Code) } }) t.Run("exhaust B", func(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) if err := lim.RegisterBucket("demo", "bucket-a", Bucket{ Max: 100, Decay: time.Minute, Key: func(*http.Request) string { return "a-fresh" }, }); err != nil { t.Fatal(err) } if err := lim.RegisterBucket("demo", "bucket-b", Bucket{ Max: 1, Decay: time.Minute, Key: func(*http.Request) string { return "b-only" }, }); err != nil { t.Fatal(err) } h := stackThrottle(lim, okHandler, "bucket-a", "bucket-b") req := httptest.NewRequest(http.MethodGet, "/x", nil) first := httptest.NewRecorder() h.ServeHTTP(first, req) if first.Code != http.StatusOK { t.Fatalf("first = %d", first.Code) } rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusTooManyRequests { t.Fatalf("exhausted B status = %d", rec.Code) } }) } func TestFixedWindowLimiterInlineThrottleKeys(t *testing.T) { okHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) }) t.Run("principals differ", func(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) h := lim.Middleware("1,1")(okHandler) reqA := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) reqA.RemoteAddr = "192.0.2.1:1" reqA = reqA.WithContext(bouncer.WithUser(reqA.Context(), &bouncer.Principal{ID: 1})) reqB := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) reqB.RemoteAddr = "192.0.2.1:1" reqB = reqB.WithContext(bouncer.WithUser(reqB.Context(), &bouncer.Principal{ID: 2})) a1 := httptest.NewRecorder() h.ServeHTTP(a1, reqA) a2 := httptest.NewRecorder() h.ServeHTTP(a2, reqA) if a2.Code != http.StatusTooManyRequests { t.Fatalf("user 1 second status = %d", a2.Code) } b1 := httptest.NewRecorder() h.ServeHTTP(b1, reqB) if b1.Code != http.StatusOK { t.Fatalf("user 2 should have a distinct key, status = %d", b1.Code) } }) 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://first.example/x", nil) req1.RemoteAddr = "192.0.2.1:1" req2 := httptest.NewRequest(http.MethodGet, "http://second.example/x", nil) req2.RemoteAddr = "192.0.2.1:9" first := httptest.NewRecorder() h.ServeHTTP(first, req1) if first.Code != http.StatusOK { t.Fatalf("first = %d", first.Code) } second := httptest.NewRecorder() h.ServeHTTP(second, req2) if second.Code != http.StatusTooManyRequests { 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) } }) } func TestFixedWindowLimiterUnknownBucketFailsAssemble(t *testing.T) { p := assemblePlugin{id: "golem15.demo", use: []string{"throttle:missing"}} _, err := Assemble(backpack.New(nil), []party.Plugin{p}) if err == nil || !strings.Contains(err.Error(), "missing") { t.Fatalf("want unknown throttle in error, got %v", err) } } func TestFixedWindowLimiterDuplicateBucket(t *testing.T) { lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) b := Bucket{Max: 1, Decay: time.Minute, Key: func(*http.Request) string { return "k" }} if err := lim.RegisterBucket("one", "shared", b); err != nil { t.Fatal(err) } err := lim.RegisterBucket("two", "shared", b) if err == nil || !strings.Contains(err.Error(), "shared") || !strings.Contains(err.Error(), "one") { t.Fatalf("got %v", err) } } func stackThrottle(lim *FixedWindowLimiter, next http.Handler, names ...string) http.Handler { h := next for i := len(names) - 1; i >= 0; i-- { h = lim.Middleware(names[i])(h) } return h } var _ pact.Router = (*Router)(nil) func TestRegisterBucketRejectsInvalidDefinitions(t *testing.T) { l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil) key := func(*http.Request) string { return "k" } if err := l.RegisterBucket("p", "n", Bucket{Max: 1, Decay: 0, Key: key}); err == nil { t.Fatal("zero decay accepted") } if err := l.RegisterBucket("p", "n", Bucket{Max: 1, Decay: time.Minute}); err == nil { t.Fatal("nil key accepted") } if err := l.ValidateThrottle("1,9223372036854775807"); err == nil { t.Fatal("overflowing minutes accepted") } } func TestRegisterBucketRejectsInvalid(t *testing.T) { key := func(*http.Request) string { return "k" } cases := []struct { name string store Store b Bucket }{ {"nil key", NewMemoryStore(time.Minute), Bucket{Max: 1, Decay: time.Minute}}, {"max zero", NewMemoryStore(time.Minute), Bucket{Max: 0, Decay: time.Minute, Key: key}}, {"max negative", NewMemoryStore(time.Minute), Bucket{Max: -1, Decay: time.Minute, Key: key}}, {"decay zero", NewMemoryStore(time.Minute), Bucket{Max: 1, Key: key}}, {"decay negative", NewMemoryStore(time.Minute), Bucket{Max: 1, Decay: -time.Second, Key: key}}, {"nil store", nil, Bucket{Max: 1, Decay: time.Minute, Key: key}}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { l := NewFixedWindowLimiter(tc.store, nil) err := l.RegisterBucket("golem15.p", "bkt", tc.b) if err == nil || !strings.Contains(err.Error(), "golem15.p") || !strings.Contains(err.Error(), "bkt") { t.Fatalf("err = %v", err) } }) } } func TestValidateThrottleRejectsOverflowAndNilStore(t *testing.T) { l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil) for _, p := range []string{"1,9223372036854775807", "0,1", "1,0", "-1,1", "x,y", "nope"} { if err := l.ValidateThrottle(p); err == nil { t.Errorf("%q accepted", p) } } if err := l.ValidateThrottle("5,1"); err != nil { t.Fatal(err) } if err := NewFixedWindowLimiter(nil, nil).ValidateThrottle("5,1"); err == nil { t.Fatal("nil store accepted") } var nilLim *FixedWindowLimiter if err := nilLim.ValidateThrottle("5,1"); err == nil { t.Fatal("nil limiter accepted") } } func TestMiddlewareFailsClosed(t *testing.T) { var nilLim *FixedWindowLimiter cases := map[string]*FixedWindowLimiter{ "unknown bucket": NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil), "nil store": NewFixedWindowLimiter(nil, nil), "nil limiter": nilLim, } params := map[string]string{"unknown bucket": "missing", "nil store": "5,1", "nil limiter": "5,1"} for name, l := range cases { t.Run(name, func(t *testing.T) { called := false h := l.Middleware(params[name])(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { called = true w.WriteHeader(http.StatusNoContent) })) rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) if called || rec.Code != http.StatusInternalServerError { t.Fatalf("called=%v code=%d", called, rec.Code) } }) } } func TestRegisterBucketRejectsNilLimiterEmptyNameAndDuplicate(t *testing.T) { key := func(*http.Request) string { return "k" } good := Bucket{Max: 1, Decay: time.Minute, Key: key} var nilLim *FixedWindowLimiter if err := nilLim.RegisterBucket("p", "n", good); err == nil { t.Fatal("nil limiter accepted") } l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil) if err := l.RegisterBucket("p", "", good); err == nil { t.Fatal("empty name accepted") } if err := l.RegisterBucket("p", "n", good); err != nil { t.Fatal(err) } if err := l.RegisterBucket("q", "n", good); err == nil || !strings.Contains(err.Error(), "p") { t.Fatalf("duplicate: %v", err) } }