package ratelimit import ( "net/http" "net/http/httptest" "testing" "time" ) func TestMiddleware(t *testing.T) { tests := map[string]struct { config Config requests int expectedCode int }{ "allows requests within limit": { config: Config{ RPS: 2, Burst: 2, TTL: time.Minute, }, requests: 1, expectedCode: http.StatusOK, }, "blocks requests over limit": { config: Config{ RPS: 1, Burst: 1, TTL: time.Minute, }, requests: 2, expectedCode: http.StatusTooManyRequests, }, "allows burst requests": { config: Config{ RPS: 1, Burst: 3, TTL: time.Minute, }, requests: 3, expectedCode: http.StatusOK, }, } for name, tt := range tests { t.Run(name, func(t *testing.T) { handler := Middleware(tt.config) nextHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }) var lastCode int for range tt.requests { w := httptest.NewRecorder() r := httptest.NewRequest(http.MethodGet, "/", nil) handler(nextHandler).ServeHTTP(w, r) lastCode = w.Code } if lastCode != tt.expectedCode { t.Errorf("expected status %d, got %d", tt.expectedCode, lastCode) } }) } } func TestRateLimiter_getVisitor(t *testing.T) { limiter := newLimiter(10, 20, time.Second) ip := visitorIP("127.0.0.1") v := limiter.getVisitor(ip) if v == nil { t.Fatal("expected non-nil limiter") } vAgain := limiter.getVisitor(ip) if v != vAgain { t.Fatal("expected the same limiter for the same IP") } if want := 1; len(limiter.visitors) != want { t.Fatalf("expected %d visitor, got %d", want, len(limiter.visitors)) } } func TestRateLimiter_cleanupVisitors(t *testing.T) { limiter := newLimiter(10, 20, time.Millisecond) limiter.getVisitor("192.168.9.1") if want := 1; len(limiter.visitors) != want { t.Fatalf("expected %d visitor, got %d", want, len(limiter.visitors)) } time.Sleep(5 * time.Millisecond) limiter.cleanupVisitors() if want := 0; len(limiter.visitors) != want { t.Fatalf("expected %d visitors after cleanup, got %d", want, len(limiter.visitors)) } } func TestRateLimiter_differentIPs(t *testing.T) { limiter := newLimiter(10, 20, time.Second) ip1 := limiter.getVisitor("1.1.1.1") ip2 := limiter.getVisitor("2.2.2.2") if ip1 == ip2 { t.Fatal("expected different limiters for different IPs") } if want := 2; len(limiter.visitors) != want { t.Fatalf("expected %d visitors, got %d", want, len(limiter.visitors)) } } func TestGetIP(t *testing.T) { t.Run("uses X-Forwarded-For", func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, "/", nil) r.Header.Set("X-Forwarded-For", "203.0.113.1") if got := getIP(r); got != "203.0.113.1" { t.Errorf("expected %q, got %q", "203.0.113.1", got) } }) t.Run("uses first IP from X-Forwarded-For list", func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, "/", nil) r.Header.Set("X-Forwarded-For", "203.0.113.1, 198.51.100.2, 192.0.2.3") if got := getIP(r); got != "203.0.113.1" { t.Errorf("expected %q, got %q", "203.0.113.1", got) } }) t.Run("uses X-Real-IP when no X-Forwarded-For", func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, "/", nil) r.Header.Set("X-Real-IP", "10.0.0.1") if got := getIP(r); got != "10.0.0.1" { t.Errorf("expected %q, got %q", "10.0.0.1", got) } }) t.Run("falls back to RemoteAddr", func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, "/", nil) r.RemoteAddr = "192.168.1.1:12345" if got := getIP(r); got != "192.168.1.1" { t.Errorf("expected %q, got %q", "192.168.1.1", got) } }) }