package ratelimit import ( "net/http" "testing" "time" ) func TestAllowsUpToTheLimitThenBlocks(t *testing.T) { l := New(3, time.Minute) for i := range 3 { if !l.Allow("1.1.1.1") { t.Fatalf("call %d should be allowed", i) } } if l.Allow("1.1.1.1") { t.Fatal("call past the limit should be blocked") } } func TestBucketsArePerKey(t *testing.T) { l := New(1, time.Minute) if !l.Allow("3.3.3.1") || l.Allow("3.3.3.1") { t.Fatal("first key not limited correctly") } if !l.Allow("3.3.3.2") { t.Fatal("second key should have its own bucket") } } func TestLimitersAreIndependent(t *testing.T) { login := New(1, time.Minute) comment := New(1, time.Minute) if !login.Allow("2.2.2.2") || login.Allow("2.2.2.2") { t.Fatal("login limiter not limited correctly") } if !comment.Allow("2.2.2.2") { t.Fatal("comment limiter should be independent") } } func TestWindowResets(t *testing.T) { l := New(1, 20*time.Millisecond) if !l.Allow("4.4.4.4") || l.Allow("4.4.4.4") { t.Fatal("limit not enforced") } time.Sleep(40 * time.Millisecond) if !l.Allow("4.4.4.4") { t.Fatal("window should have reset") } } func TestCountingStaysCorrectWithManyKeys(t *testing.T) { l := New(5, time.Minute) for i := range 1200 { l.Allow("5.5.5.5-" + string(rune('a'+i%26)) + string(rune('a'+i/26))) } for i := range 5 { if !l.Allow("5.5.5.5") { t.Fatalf("call %d should be allowed", i) } } if l.Allow("5.5.5.5") { t.Fatal("call past the limit should be blocked") } } func TestClientIP(t *testing.T) { r, err := http.NewRequest("GET", "http://x/", nil) if err != nil { t.Fatal(err) } r.Header.Set("X-Forwarded-For", "9.9.9.9, 10.0.0.1") r.RemoteAddr = "10.0.0.1:5555" if got := ClientIP(r, true); got != "9.9.9.9" { t.Errorf("trusted proxy: got %q", got) } if got := ClientIP(r, false); got != "10.0.0.1" { t.Errorf("untrusted: got %q", got) } } func TestClientIPFallbackAndIPv6Prefix(t *testing.T) { r, err := http.NewRequest("GET", "http://x/", nil) if err != nil { t.Fatal(err) } r.RemoteAddr = "10.0.0.1:5555" if got := ClientIP(r, true); got != "10.0.0.1" { t.Errorf("trusted proxy without header: got %q", got) } r.RemoteAddr = "[2001:db8:1:2:aaaa::1]:5555" a := ClientIP(r, false) r.RemoteAddr = "[2001:db8:1:2:bbbb::9]:5555" if b := ClientIP(r, false); a != b || a != "2001:db8:1:2::/64" { t.Errorf("same /64 got %q and %q", a, b) } r.RemoteAddr = "[::ffff:10.0.0.2]:5555" if got := ClientIP(r, false); got != "10.0.0.2" { t.Errorf("mapped IPv4: got %q", got) } r.RemoteAddr = "" if got := ClientIP(r, true); got == "" { t.Error("key must never be empty") } }