package middleware_test import ( "net/http" "net/http/httptest" "testing" "github.com/gin-gonic/gin" "github.com/wangjia/pay/internal/middleware" ) func newTestEngine(rl *middleware.IPRateLimiter) *gin.Engine { gin.SetMode(gin.TestMode) r := gin.New() r.GET("/x", rl.Gin(), func(c *gin.Context) { c.Status(http.StatusOK) }) return r } // burst(=RequestsPerMin)个请求应放行,第 burst+1 个应 429。 func TestIPRateLimiterBurstThenBlocks(t *testing.T) { rl := middleware.NewIPRateLimiter(3) r := newTestEngine(rl) for i := 0; i < 3; i++ { w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/x", nil)) if w.Code != http.StatusOK { t.Fatalf("第 %d 次(burst 内)应 200, got %d", i+1, w.Code) } } w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/x", nil)) if w.Code != http.StatusTooManyRequests { t.Fatalf("第 4 次(超 burst)应 429, got %d", w.Code) } } // 不同 IP 各自独立计桶,互不影响。 func TestIPRateLimiterPerIPIsolated(t *testing.T) { rl := middleware.NewIPRateLimiter(1) r := newTestEngine(rl) req1 := httptest.NewRequest(http.MethodGet, "/x", nil) req1.RemoteAddr = "10.0.0.1:1234" w1 := httptest.NewRecorder() r.ServeHTTP(w1, req1) if w1.Code != http.StatusOK { t.Fatalf("IP1 第 1 次应 200, got %d", w1.Code) } req2 := httptest.NewRequest(http.MethodGet, "/x", nil) req2.RemoteAddr = "10.0.0.2:1234" w2 := httptest.NewRecorder() r.ServeHTTP(w2, req2) if w2.Code != http.StatusOK { t.Fatalf("IP2(独立桶)第 1 次应 200, got %d", w2.Code) } // IP1 burst=1 已耗尽,第 2 次应 429。 req3 := httptest.NewRequest(http.MethodGet, "/x", nil) req3.RemoteAddr = "10.0.0.1:1234" w3 := httptest.NewRecorder() r.ServeHTTP(w3, req3) if w3.Code != http.StatusTooManyRequests { t.Fatalf("IP1 第 2 次(超 burst)应 429, got %d", w3.Code) } } // perMinute<=0 兜底成默认 30/min(config.RateLimitConfig 零值语义:"零值=默认开")。 func TestNewIPRateLimiterZeroDefaultsTo30(t *testing.T) { rl := middleware.NewIPRateLimiter(0) r := newTestEngine(rl) for i := 0; i < 30; i++ { w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/x", nil)) if w.Code != http.StatusOK { t.Fatalf("零值兜底 30/min:第 %d 次应 200, got %d", i+1, w.Code) } } w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/x", nil)) if w.Code != http.StatusTooManyRequests { t.Fatalf("第 31 次(超默认 30/min)应 429, got %d", w.Code) } }