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 TestDialControlRejectsUnsafeIPv6Transitions(t *testing.T) { tests := []struct { name string ip string }{ {name: "nat64 well-known loopback", ip: "64:ff9b::7f00:1"}, {name: "nat64 well-known rfc1918", ip: "64:ff9b::a00:1"}, {name: "nat64 well-known metadata", ip: "64:ff9b::a9fe:a9fe"}, {name: "nat64 local-use loopback", ip: "64:ff9b:1:7f00:0:100::"}, {name: "nat64 local-use rfc1918", ip: "64:ff9b:1:a00:0:100::"}, {name: "nat64 local-use metadata", ip: "64:ff9b:1:a9fe:a9:fe00::"}, {name: "6to4 loopback", ip: "2002:7f00:1::"}, {name: "6to4 rfc1918", ip: "2002:a00:1::"}, {name: "6to4 metadata", ip: "2002:a9fe:a9fe::"}, } control := dialControl(Policy{Mode: PublicOnlyMode}) for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { err := control("tcp6", "["+tt.ip+"]:443", nil) if !errors.Is(err, errPrivateIP) { t.Fatalf("dialControl(%s) error = %v, want errPrivateIP", tt.ip, err) } if got := mapTransportError(err).Reason; got != ReasonPrivateIP { t.Fatalf("mapTransportError(%s) reason = %q, want %q", tt.ip, got, 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 } func TestDialControlRejectsZonedAndSpecialUse(t *testing.T) { ctl := dialControl(Policy{}) for _, a := range []string{"[fe80::1%eth0]:443", "198.18.0.1:443"} { if err := ctl("tcp", a, nil); !errors.Is(err, errPrivateIP) { t.Errorf("%s: got %v", a, err) } } }