package api import ( "context" "net/http" "net/http/httptest" "strings" "testing" "time" ) func TestParseRateLimit(t *testing.T) { tests := []struct { in string limit int window time.Duration }{ {"", defaultRateLimitCount, defaultRateLimitWindow}, {"30/1h", 30, time.Hour}, {"10/10m", 10, 10 * time.Minute}, {"5 / 30s", 5, 30 * time.Second}, {"bogus", defaultRateLimitCount, defaultRateLimitWindow}, {"0/1h", defaultRateLimitCount, defaultRateLimitWindow}, {"-3/1h", defaultRateLimitCount, defaultRateLimitWindow}, {"10/-1h", defaultRateLimitCount, defaultRateLimitWindow}, {"10/soon", defaultRateLimitCount, defaultRateLimitWindow}, {"/1h", defaultRateLimitCount, defaultRateLimitWindow}, } for _, tc := range tests { limit, window := parseRateLimit(tc.in) if limit != tc.limit || window != tc.window { t.Errorf("parseRateLimit(%q) = %d, %s; want %d, %s", tc.in, limit, window, tc.limit, tc.window) } } } func TestNewHandlerReadsRateLimitEnv(t *testing.T) { t.Setenv("EXPLOREDNS_RATE_LIMIT", "5/10m") ctx, cancel := context.WithCancel(context.Background()) defer cancel() h := newHandler(ctx) if h.limiter.limit != 5 || h.limiter.window != 10*time.Minute { t.Fatalf("limiter = %d/%s, want 5/10m", h.limiter.limit, h.limiter.window) } } // postTraverse sends POST /api/traverse with an empty JSON body so requests // that pass the rate limiter fail validation (400) instead of spawning a // real traversal. remoteAddr and headers shape the client identity. func postTraverse(t *testing.T, h *Handler, remoteAddr string, headers map[string]string) *httptest.ResponseRecorder { t.Helper() req := httptest.NewRequest(http.MethodPost, "/api/traverse", strings.NewReader(`{}`)) req.RemoteAddr = remoteAddr for k, v := range headers { req.Header.Set(k, v) } rec := httptest.NewRecorder() h.mux.ServeHTTP(rec, req) return rec } func TestRateLimit_OverLimitReturns429(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() h := newHandler(ctx) h.limiter = newRateLimiter(2, time.Hour) hdr := map[string]string{"X-Forwarded-For": "203.0.113.9"} for i := 0; i < 2; i++ { if rec := postTraverse(t, h, "10.0.0.1:1234", hdr); rec.Code != http.StatusBadRequest { t.Fatalf("request %d: want 400 (under limit), got %d: %s", i, rec.Code, rec.Body.String()) } } rec := postTraverse(t, h, "10.0.0.1:1234", hdr) if rec.Code != http.StatusTooManyRequests { t.Fatalf("want 429 over limit, got %d: %s", rec.Code, rec.Body.String()) } if body := rec.Body.String(); !strings.Contains(body, "2 requests per 1h0m0s") { t.Fatalf("429 body should name the limit, got %s", body) } // A distinct client IP has its own bucket and is unaffected. other := map[string]string{"X-Forwarded-For": "203.0.113.10"} if rec := postTraverse(t, h, "10.0.0.1:1234", other); rec.Code != http.StatusBadRequest { t.Fatalf("distinct IP: want 400, got %d: %s", rec.Code, rec.Body.String()) } } func TestRateLimit_RefillsContinuously(t *testing.T) { rl := newRateLimiter(2, time.Second) now := time.Now() rl.now = func() time.Time { return now } if !rl.allow("a") || !rl.allow("a") { t.Fatal("first two requests should be allowed") } if rl.allow("a") { t.Fatal("third request should be denied") } // Half a window refills half the bucket: one token. now = now.Add(500 * time.Millisecond) if !rl.allow("a") { t.Fatal("request after refill should be allowed") } if rl.allow("a") { t.Fatal("bucket should hold only the refilled token") } } func TestRateLimit_SweepDropsIdleBuckets(t *testing.T) { rl := newRateLimiter(1, time.Minute) now := time.Now() rl.now = func() time.Time { return now } rl.allow("stale") now = now.Add(2 * time.Minute) rl.allow("fresh") rl.sweep() rl.mu.Lock() defer rl.mu.Unlock() if _, ok := rl.buckets["stale"]; ok { t.Fatal("idle bucket should have been swept") } if _, ok := rl.buckets["fresh"]; !ok { t.Fatal("active bucket should survive the sweep") } } func TestRateLimit_LocalhostExempt(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() h := newHandler(ctx) h.limiter = newRateLimiter(1, time.Hour) for _, addr := range []string{"127.0.0.1:5555", "[::1]:5555"} { for i := 0; i < 3; i++ { rec := postTraverse(t, h, addr, nil) if rec.Code != http.StatusBadRequest { t.Fatalf("%s request %d: localhost should be exempt, got %d: %s", addr, i, rec.Code, rec.Body.String()) } } } } // TestRateLimit_ProxiedLocalhostNotExempt verifies that a request arriving // from a local proxy (RemoteAddr loopback) is still limited by the real // client IP carried in Fly-Client-IP / X-Forwarded-For. func TestRateLimit_ProxiedLocalhostNotExempt(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() h := newHandler(ctx) h.limiter = newRateLimiter(1, time.Hour) hdr := map[string]string{"Fly-Client-IP": "198.51.100.4"} if rec := postTraverse(t, h, "127.0.0.1:5555", hdr); rec.Code != http.StatusBadRequest { t.Fatalf("first proxied request: want 400, got %d", rec.Code) } if rec := postTraverse(t, h, "127.0.0.1:5555", hdr); rec.Code != http.StatusTooManyRequests { t.Fatalf("second proxied request: want 429, got %d", rec.Code) } // The same proxy forwarding a different client is unaffected. other := map[string]string{"Fly-Client-IP": "198.51.100.5"} if rec := postTraverse(t, h, "127.0.0.1:5555", other); rec.Code != http.StatusBadRequest { t.Fatalf("other client via proxy: want 400, got %d", rec.Code) } }