package ipratelimiter import ( "testing" "golang.org/x/time/rate" ) func TestNewIPRateLimiter(t *testing.T) { l := NewIPRateLimiter(rate.Limit(100), 100) if l == nil { t.Fatal("NewIPRateLimiter 返回 nil") } lim := l.GetLimiter("10.1.1.1") if lim == nil { t.Fatal("GetLimiter 不应返回 nil") } } func TestGetLimiterSameInstance(t *testing.T) { l := NewIPRateLimiter(rate.Limit(100), 100) lim1 := l.GetLimiter("192.168.1.1") lim2 := l.GetLimiter("192.168.1.1") if lim1 != lim2 { t.Error("同一 IP 的 GetLimiter 应返回相同实例") } } func TestGetLimiterDifferentIP(t *testing.T) { l := NewIPRateLimiter(rate.Limit(100), 100) lim1 := l.GetLimiter("192.168.1.1") lim2 := l.GetLimiter("10.0.0.1") if lim1 == lim2 { t.Error("不同 IP 的 GetLimiter 应返回不同实例") } } func TestAddIP(t *testing.T) { l := NewIPRateLimiter(rate.Limit(100), 100) lim1 := l.AddIP("172.16.0.1") lim2 := l.GetLimiter("172.16.0.1") if lim1 != lim2 { t.Error("AddIP 后 GetLimiter 应返回同一实例") } } func TestAllow(t *testing.T) { // 大速率:短时间内必然放行 l := NewIPRateLimiter(rate.Limit(1000), 1000) lim := l.GetLimiter("192.168.1.100") for i := 0; i < 10; i++ { if !lim.Allow() { t.Fatalf("第 %d 次请求被拒绝", i+1) } } }