package originauth import ( "net/http" "net/http/httptest" "testing" "github.com/go-chi/chi/v5" ) func testMW(t *testing.T) *Middleware { t.Helper() mw, err := New(Config{ Current: "secret-current", Previous: "secret-previous", AllowedCIDRs: []string{"203.0.113.0/24", "2001:db8::/32"}, }) if err != nil { t.Fatalf("New: %v", err) } return mw } func do(t *testing.T, mw *Middleware, remoteAddr, headerVal string) int { t.Helper() h := mw.Handler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("ok")) })) req := httptest.NewRequest(http.MethodGet, "/v1/nodes", nil) req.RemoteAddr = remoteAddr if headerVal != "" { req.Header.Set(DefaultHeader, headerVal) } rec := httptest.NewRecorder() h.ServeHTTP(rec, req) return rec.Code } func TestAllowedIPCurrentValue(t *testing.T) { mw := testMW(t) if code := do(t, mw, "203.0.113.5:443", "secret-current"); code != http.StatusOK { t.Fatalf("want 200, got %d", code) } } func TestAllowedIPPreviousValueDuringRotation(t *testing.T) { mw := testMW(t) if code := do(t, mw, "203.0.113.5:443", "secret-previous"); code != http.StatusOK { t.Fatalf("rotation window: want 200, got %d", code) } } func TestAllowedIPWrongValue(t *testing.T) { mw := testMW(t) if code := do(t, mw, "203.0.113.5:443", "nope"); code != http.StatusForbidden { t.Fatalf("want 403, got %d", code) } } func TestAllowedIPMissingValue(t *testing.T) { mw := testMW(t) if code := do(t, mw, "203.0.113.5:443", ""); code != http.StatusForbidden { t.Fatalf("want 403, got %d", code) } } func TestDirectConnectBypassRejected(t *testing.T) { mw := testMW(t) // Correct header but a source IP outside the CDN range = someone hitting the // origin directly. Must be 403. if code := do(t, mw, "198.51.100.7:55000", "secret-current"); code != http.StatusForbidden { t.Fatalf("direct connect must be 403, got %d", code) } } func TestIPv6Allowed(t *testing.T) { mw := testMW(t) if code := do(t, mw, "[2001:db8::1]:443", "secret-current"); code != http.StatusOK { t.Fatalf("want 200 for allowed v6, got %d", code) } } func TestForgedXFFDoesNotBypass(t *testing.T) { mw := testMW(t) h := mw.Handler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) })) req := httptest.NewRequest(http.MethodGet, "/v1/nodes", nil) req.RemoteAddr = "198.51.100.7:55000" // real peer: not a CDN range req.Header.Set("X-Forwarded-For", "203.0.113.5") req.Header.Set(DefaultHeader, "secret-current") rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusForbidden { t.Fatalf("forged XFF must not bypass, got %d", rec.Code) } } func TestNewValidation(t *testing.T) { if _, err := New(Config{AllowedCIDRs: []string{"203.0.113.0/24"}}); err == nil { t.Fatal("want error: missing Current") } if _, err := New(Config{Current: "x"}); err == nil { t.Fatal("want error: missing CIDRs") } if _, err := New(Config{Current: "x", AllowedCIDRs: []string{"not-a-cidr"}}); err == nil { t.Fatal("want error: bad CIDR") } } func TestIntegrationWithChiRouter(t *testing.T) { mw := testMW(t) r := chi.NewRouter() r.Use(mw.Handler) r.Get("/v1/nodes", func(w http.ResponseWriter, r *http.Request) { _, _ = w.Write([]byte("nodes")) }) srv := httptest.NewServer(r) defer srv.Close() // Through-the-stack request with a non-CDN client IP (httptest loopback) // must be rejected even with the correct header. req, _ := http.NewRequest(http.MethodGet, srv.URL+"/v1/nodes", nil) req.Header.Set(DefaultHeader, "secret-current") resp, err := http.DefaultClient.Do(req) if err != nil { t.Fatal(err) } defer resp.Body.Close() if resp.StatusCode != http.StatusForbidden { t.Fatalf("loopback (non-CDN) should be 403, got %d", resp.StatusCode) } }