From bae9de878846361ee78dffc7afdcfe0452d34e29 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Sat, 19 Sep 2026 19:00:01 +0200 Subject: [PATCH] test(06-04): add failing tests for SSRF-guarded Fetch - Cover host allow-list, dotted-suffix bypass, dial-time private IP - Cover no-follow redirects, streaming byte cap, and D-14 config defaults --- fetchguard/fetch_test.go | 276 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 276 insertions(+) create mode 100644 fetchguard/fetch_test.go diff --git a/fetchguard/fetch_test.go b/fetchguard/fetch_test.go new file mode 100644 index 0000000..d6b75dd --- /dev/null +++ b/fetchguard/fetch_test.go @@ -0,0 +1,276 @@ +package fetchguard + +import ( + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync/atomic" + "testing" + "time" + + "git.golem15.com/golem15/summercms/compass" +) + +func TestFetchMalformedURL(t *testing.T) { + _, err := Fetch(t.Context(), "not a url", Policy{Mode: PublicOnlyMode}, nil) + if reasonFrom(t, err) != ReasonInvalidURL { + t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonInvalidURL) + } +} + +func TestFetchHTTPSchemeRejectedWithoutIO(t *testing.T) { + var hits atomic.Int64 + srv := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + hits.Add(1) + })) + t.Cleanup(srv.Close) + + _, err := Fetch(t.Context(), srv.URL, Policy{Mode: PublicOnlyMode}, nil) + if reasonFrom(t, err) != ReasonScheme { + t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonScheme) + } + if hits.Load() != 0 { + t.Fatal("http URL must not cause network I/O") + } +} + +func TestFetchAllowHostsRejectsUnknownHostBeforeDial(t *testing.T) { + var hits atomic.Int64 + srv := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + hits.Add(1) + })) + t.Cleanup(srv.Close) + + _, err := Fetch(t.Context(), srv.URL, Policy{ + Mode: AllowHostsMode, + AllowHosts: []string{"discogs.com"}, + }, nil) + if reasonFrom(t, err) != ReasonInvalidURL { + t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonInvalidURL) + } + if hits.Load() != 0 { + t.Fatal("host allow-list miss must not dial") + } +} + +func TestFetchAllowHostsRejectsDottedSuffixBypass(t *testing.T) { + _, err := Fetch(t.Context(), "https://evil-discogs.com/cover.jpg", Policy{ + Mode: AllowHostsMode, + AllowHosts: []string{"discogs.com"}, + }, nil) + if reasonFrom(t, err) != ReasonInvalidURL { + t.Fatalf("evil-discogs.com must not match discogs.com, reason = %q", reasonFrom(t, err)) + } +} + +func TestFetchPrivateIPBlockedInBothModes(t *testing.T) { + srv := httptest.NewTLSServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + t.Error("handler must not run for a private dial") + })) + t.Cleanup(srv.Close) + + t.Run("AllowHostsMode", func(t *testing.T) { + _, err := Fetch(t.Context(), srv.URL, Policy{ + Mode: AllowHostsMode, + AllowHosts: []string{"127.0.0.1"}, + Timeout: 2 * time.Second, + MaxBytes: 1024, + }, nil) + if reasonFrom(t, err) != ReasonPrivateIP { + t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonPrivateIP) + } + }) + t.Run("PublicOnlyMode", func(t *testing.T) { + _, err := Fetch(t.Context(), srv.URL, Policy{ + Mode: PublicOnlyMode, + Timeout: 2 * time.Second, + MaxBytes: 1024, + }, nil) + if reasonFrom(t, err) != ReasonPrivateIP { + t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonPrivateIP) + } + }) +} + +func TestFetchDoesNotFollowRedirect(t *testing.T) { + var followed atomic.Bool + mux := http.NewServeMux() + mux.HandleFunc("/target", func(w http.ResponseWriter, r *http.Request) { + followed.Store(true) + w.WriteHeader(http.StatusOK) + }) + mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Location", "/target") + w.WriteHeader(http.StatusFound) + }) + srv := httptest.NewTLSServer(mux) + t.Cleanup(srv.Close) + + res, err := Fetch(t.Context(), srv.URL, withTestLoopback(srv, Policy{ + Mode: PublicOnlyMode, + MaxBytes: 1024, + Timeout: 5 * time.Second, + }), nil) + if err != nil { + t.Fatalf("Fetch: %v", err) + } + if res.StatusCode != http.StatusFound { + t.Fatalf("status = %d, want %d (3xx returned, not followed)", res.StatusCode, http.StatusFound) + } + if followed.Load() { + t.Fatal("redirect must not be followed") + } +} + +func TestFetchTooLargeIsStreaming(t *testing.T) { + const maxBytes int64 = 256 + var written atomic.Int64 + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/octet-stream") + w.WriteHeader(http.StatusOK) + buf := make([]byte, 64) + for written.Load() < 8<<20 { + n, err := w.Write(buf) + written.Add(int64(n)) + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + if err != nil { + return + } + } + })) + t.Cleanup(srv.Close) + + _, err := Fetch(t.Context(), srv.URL, withTestLoopback(srv, Policy{ + Mode: PublicOnlyMode, + MaxBytes: maxBytes, + Timeout: 5 * time.Second, + }), nil) + if reasonFrom(t, err) != ReasonTooLarge { + t.Fatalf("reason = %q, want %s", reasonFrom(t, err), ReasonTooLarge) + } + // TCP/HTTP buffering can write a little past maxBytes+1; the client must + // not have pulled an unbounded body first (the handler would hit 8MiB). + if got := written.Load(); got > 64<<10 { + t.Fatalf("server wrote %d bytes, client appears to have buffered unbounded body", got) + } + if written.Load() < maxBytes+1 { + t.Fatalf("server wrote %d bytes, want at least maxBytes+1 so the cap was hit", written.Load()) + } +} + +func TestFetchZeroPolicyUsesDefaults(t *testing.T) { + max, timeout := Defaults() + if max != 10*1024*1024 { + t.Fatalf("Defaults maxBytes = %d, want 10MiB", max) + } + if timeout != 10*time.Second { + t.Fatalf("Defaults timeout = %s, want 10s", timeout) + } + + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/plain") + io.WriteString(w, "ok") + })) + t.Cleanup(srv.Close) + + res, err := Fetch(t.Context(), srv.URL, withTestLoopback(srv, Policy{Mode: PublicOnlyMode}), nil) + if err != nil { + t.Fatalf("Fetch with zero MaxBytes/Timeout and nil cfg: %v", err) + } + if string(res.Body) != "ok" { + t.Fatalf("body = %q, want ok", res.Body) + } +} + +func TestDefaultsFromConfigExplicitZeroIsError(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte("fetch:\n max_bytes: 0\n timeout_seconds: 10\n"), 0o644); err != nil { + t.Fatal(err) + } + cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}}) + if err != nil { + t.Fatalf("Open: %v", err) + } + if _, ok := cfg.Lookup("http.fetch.max_bytes"); !ok { + t.Fatal("expected http.fetch.max_bytes to be present in loaded YAML") + } + _, _, err = DefaultsFromConfig(cfg) + if err == nil { + t.Fatal("explicit http.fetch.max_bytes: 0 must be an error, not a silent fallback") + } +} + +func TestDefaultsFromConfigExplicitNegativeIsError(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte("fetch:\n max_bytes: 10485760\n timeout_seconds: -1\n"), 0o644); err != nil { + t.Fatal(err) + } + cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}}) + if err != nil { + t.Fatalf("Open: %v", err) + } + _, _, err = DefaultsFromConfig(cfg) + if err == nil { + t.Fatal("explicit http.fetch.timeout_seconds: -1 must be an error") + } +} + +func TestDefaultsFromConfigAbsentFallsBack(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: t\n"), 0o644); err != nil { + t.Fatal(err) + } + cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}}) + if err != nil { + t.Fatalf("Open: %v", err) + } + max, timeout, err := DefaultsFromConfig(cfg) + if err != nil { + t.Fatalf("DefaultsFromConfig: %v", err) + } + wantMax, wantTO := Defaults() + if max != wantMax || timeout != wantTO { + t.Fatalf("got %d %s, want Defaults() %d %s", max, timeout, wantMax, wantTO) + } +} + +func TestErrorUnwrap(t *testing.T) { + inner := errors.New("dial tcp") + err := &Error{Reason: ReasonNetworkError, Err: inner} + if err.Error() == "" { + t.Fatal("Error() must be non-empty") + } + if !errors.Is(err, inner) { + t.Fatal("Unwrap must expose the inner error") + } +} + +func reasonFrom(t *testing.T, err error) Reason { + t.Helper() + if err == nil { + t.Fatal("expected error") + } + var fe *Error + if !errors.As(err, &fe) { + t.Fatalf("err = %v (%T), want *Error", err, err) + } + return fe.Reason +} + +// withTestLoopback trusts the httptest TLS cert and skips the reserved-IP +// dial check so redirect/byte-cap tests can exercise a real listener on +// 127.0.0.1. Production Policy values leave both fields unset. +func withTestLoopback(srv *httptest.Server, p Policy) Policy { + tr, ok := srv.Client().Transport.(*http.Transport) + if !ok { + panic("httptest client transport is not *http.Transport") + } + p.tlsConfig = tr.TLSClientConfig + p.skipReservedCheck = true + return p +}