package api import ( "net/http" "net/http/httptest" "strings" "testing" "time" ) func TestRateLimiterEnforcesWindowAndBoundsKeyMemory(t *testing.T) { limiter, err := NewRateLimiter(2, time.Second, 1) if err != nil { t.Fatal(err) } start := time.Unix(1000, 0) if !limiter.Allow("player-1", start) || !limiter.Allow("player-1", start.Add(100*time.Millisecond)) { t.Fatal("allowed requests were rejected") } if limiter.Allow("player-1", start.Add(200*time.Millisecond)) { t.Fatal("request over the window limit was accepted") } if limiter.Allow("player-2", start.Add(300*time.Millisecond)) { t.Fatal("unbounded new key bypassed the memory bound") } if !limiter.Allow("player-1", start.Add(time.Second)) { t.Fatal("window did not reset at the boundary") } } func TestRateLimiterChargesCredentialAndIPDimensionsAtomically(t *testing.T) { limiter, err := NewRateLimiter(1, time.Minute, 8) if err != nil { t.Fatal(err) } start := time.Unix(1000, 0) if !limiter.AllowKeys([]string{"auth:player", "ip:one"}, start) { t.Fatal("first request was rejected") } if limiter.AllowKeys([]string{"auth:player", "ip:two"}, start) { t.Fatal("same credential bypassed the account dimension by changing IP") } if limiter.AllowKeys([]string{"auth:other", "ip:one"}, start) { t.Fatal("same IP bypassed the IP dimension by changing credential") } if !limiter.AllowKeys([]string{"auth:other", "ip:two"}, start) { t.Fatal("unrelated credential/IP pair was charged by a rejected request") } } func TestRateLimitedHTTPBoundaryReturnsGeneric429(t *testing.T) { limiter, err := NewRateLimiter(1, time.Minute, 8) if err != nil { t.Fatal(err) } service := &Service{RateLimiter: limiter, Now: func() time.Time { return time.Unix(1000, 0) }} server := httptest.NewServer(service.Handler()) defer server.Close() request, err := http.NewRequest(http.MethodGet, server.URL+"/unknown", nil) if err != nil { t.Fatal(err) } request.Header.Set("Authorization", "Bearer secret-session:secret-token") response, err := server.Client().Do(request) if err != nil { t.Fatal(err) } _ = response.Body.Close() if response.StatusCode != http.StatusNotFound { t.Fatalf("first request status = %d", response.StatusCode) } request, _ = http.NewRequest(http.MethodGet, server.URL+"/unknown", strings.NewReader("")) request.Header.Set("Authorization", "Bearer secret-session:secret-token") response, err = server.Client().Do(request) if err != nil { t.Fatal(err) } _ = response.Body.Close() if response.StatusCode != http.StatusTooManyRequests { t.Fatalf("limited request status = %d", response.StatusCode) } } func TestClientIPResolverTrustsForwardingOnlyFromConfiguredProxy(t *testing.T) { resolver, err := NewClientIPResolver("10.0.0.0/8, 2001:db8::/32") if err != nil { t.Fatal(err) } untrusted := httptest.NewRequest(http.MethodGet, "/", nil) untrusted.RemoteAddr = "203.0.113.10:1234" untrusted.Header.Set("X-Forwarded-For", "198.51.100.7") if got := requestClientIP(untrusted, resolver); got != "203.0.113.10" { t.Fatalf("untrusted peer selected forwarded IP %q", got) } trusted := httptest.NewRequest(http.MethodGet, "/", nil) trusted.RemoteAddr = "10.2.3.4:443" trusted.Header.Set("X-Forwarded-For", "198.51.100.7, 10.9.8.7") if got := requestClientIP(trusted, resolver); got != "198.51.100.7" { t.Fatalf("trusted proxy chain resolved to %q", got) } trusted.Header["X-Forwarded-For"] = []string{"192.0.2.99", "198.51.100.7, 10.9.8.7"} if got := requestClientIP(trusted, resolver); got != "198.51.100.7" { t.Fatalf("repeated forwarded headers bypassed the nearest untrusted address: %q", got) } trusted.Header.Set("X-Forwarded-For", "forged, 198.51.100.7") if got := requestClientIP(trusted, resolver); got != "10.2.3.4" { t.Fatalf("malformed forwarding did not fail closed to immediate peer: %q", got) } } func TestClientIPResolverRejectsInvalidCIDRs(t *testing.T) { if _, err := NewClientIPResolver("10.0.0.0/8,not-a-network"); err == nil { t.Fatal("invalid trusted proxy CIDR accepted") } } func TestRateLimiterSeparatesClientsBehindTrustedGateway(t *testing.T) { limiter, err := NewRateLimiter(1, time.Minute, 8) if err != nil { t.Fatal(err) } resolver, err := NewClientIPResolver("127.0.0.0/8") if err != nil { t.Fatal(err) } service := &Service{RateLimiter: limiter, ClientIPs: resolver, Now: func() time.Time { return time.Unix(1000, 0) }} server := httptest.NewServer(service.Handler()) defer server.Close() request := func(forwarded string) int { req, requestErr := http.NewRequest(http.MethodGet, server.URL+"/unknown", nil) if requestErr != nil { t.Fatal(requestErr) } req.Header.Set("X-Forwarded-For", forwarded) response, requestErr := server.Client().Do(req) if requestErr != nil { t.Fatal(requestErr) } response.Body.Close() return response.StatusCode } if got := request("198.51.100.1"); got != http.StatusNotFound { t.Fatalf("first client status = %d", got) } if got := request("198.51.100.2"); got != http.StatusNotFound { t.Fatalf("second client behind gateway status = %d", got) } if got := request("198.51.100.1"); got != http.StatusTooManyRequests { t.Fatalf("repeated first client status = %d", got) } }