package middleware import ( "net/http" "net/http/httptest" "testing" "time" "github.com/gin-gonic/gin" "fileshare/internal/cache" ) func rateLimitRouter(rl *RateLimiter, kind string) *gin.Engine { r := gin.New() r.Use(ClientIP(nil)) r.GET("/limited", rl.RequireRateLimit(kind), func(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"ok": true}) }) return r } func TestRateLimiterBlocksAfterCount(t *testing.T) { mem := cache.NewMemory() defer mem.Close() rl := NewRateLimiter(mem, map[string]LimitRule{ LimitMeta: {Count: 3, Window: time.Minute}, }) r := rateLimitRouter(rl, LimitMeta) // 前 3 次通过 for i := 0; i < 3; i++ { w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest("GET", "/limited", nil)) if w.Code != http.StatusOK { t.Fatalf("第 %d 次应通过: %d", i+1, w.Code) } } // 第 4 次 423 w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest("GET", "/limited", nil)) if w.Code != http.StatusLocked { t.Fatalf("超限应 423: %d", w.Code) } } func TestRateLimiterAddAfterSuccess(t *testing.T) { // 模拟 upload 语义:Check 放行 + handler 成功后 Add mem := cache.NewMemory() defer mem.Close() rl := NewRateLimiter(mem, map[string]LimitRule{ LimitUpload: {Count: 2, Window: time.Minute}, }) r := gin.New() r.Use(ClientIP(nil)) r.POST("/upload", func(c *gin.Context) { if allowed, _ := rl.Check(c, LimitUpload); !allowed { c.JSON(http.StatusLocked, gin.H{"err": "too many"}) return } rl.Add(c, LimitUpload) // 成功上传计数 c.JSON(http.StatusOK, gin.H{"ok": true}) }) for i := 0; i < 2; i++ { w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest("POST", "/upload", nil)) if w.Code != http.StatusOK { t.Fatalf("第 %d 次上传应通过", i+1) } } w := httptest.NewRecorder() r.ServeHTTP(w, httptest.NewRequest("POST", "/upload", nil)) if w.Code != http.StatusLocked { t.Fatalf("第 3 次上传应被拒绝: %d", w.Code) } } func TestParseTrustedProxies(t *testing.T) { nets := ParseTrustedProxies([]string{"10.0.0.0/8", "192.168.1.1", "", "bad-input"}) if len(nets) != 2 { t.Fatalf("应解析出 2 个可信网段: %d", len(nets)) } } func TestClientIPFromTrustedProxy(t *testing.T) { // 对齐参考语义:仅当直连地址可信时才解析 XFF, // 且从右往左返回第一个"非可信代理"地址(该地址即最近可信代理看到的客户端)。 nets := ParseTrustedProxies([]string{"127.0.0.0/8", "10.0.0.0/8"}) r := gin.New() r.Use(ClientIP(nets)) var seen string r.GET("/ip", func(c *gin.Context) { seen = GetClientIP(c) }) req := httptest.NewRequest("GET", "/ip", nil) req.RemoteAddr = "127.0.0.1:5000" req.Header.Set("X-Forwarded-For", "203.0.113.9, 10.0.0.1") r.ServeHTTP(httptest.NewRecorder(), req) if seen != "203.0.113.9" { t.Fatalf("多级可信代理下应取最左非可信地址: %s", seen) } // 仅直连可信:XFF 右起第一个(10.0.0.1)非可信 → 取它 nets2 := ParseTrustedProxies([]string{"127.0.0.0/8"}) r2 := gin.New() r2.Use(ClientIP(nets2)) var seen2 string r2.GET("/ip", func(c *gin.Context) { seen2 = GetClientIP(c) }) req2 := httptest.NewRequest("GET", "/ip", nil) req2.RemoteAddr = "127.0.0.1:5000" req2.Header.Set("X-Forwarded-For", "203.0.113.9, 10.0.0.1") r2.ServeHTTP(httptest.NewRecorder(), req2) if seen2 != "10.0.0.1" { t.Fatalf("右起第一个非可信地址应为 10.0.0.1: %s", seen2) } // 非可信直连:忽略伪造头 seen = "" req = httptest.NewRequest("GET", "/ip", nil) req.RemoteAddr = "8.8.8.8:1234" req.Header.Set("X-Forwarded-For", "1.2.3.4") r.ServeHTTP(httptest.NewRecorder(), req) if seen != "8.8.8.8" { t.Fatalf("非可信直连应忽略 XFF: %s", seen) } }