diff --git a/internal/httpfetcher/dial_context_internal_test.go b/internal/httpfetcher/dial_context_internal_test.go new file mode 100644 index 0000000..4553799 --- /dev/null +++ b/internal/httpfetcher/dial_context_internal_test.go @@ -0,0 +1,66 @@ +package httpfetcher + +import ( + "errors" + "net" + "testing" +) + +// TestNewUsesCheckedDialerWithoutDialContext checks that a fetcher built +// without DialContext, as pixa builds it, refuses to connect to a local +// server. +func TestNewUsesCheckedDialerWithoutDialContext(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + transport := transportOf(t, New(DefaultConfig())) + + addr := srv.Listener.Addr().String() + + _, err := transport.DialContext(testContext(t), "tcp", addr) + if !errors.Is(err, ErrSSRFBlocked) { + t.Fatalf("DialContext(%s) error = %v, want ErrSSRFBlocked", addr, err) + } +} + +// TestDialContextReplacesOnlyTheDialer checks that a fetcher built with +// DialContext connects through it, while the URL check still refuses a +// loopback URL and the redirect check a redirect to a link-local address. +func TestDialContextReplacesOnlyTheDialer(t *testing.T) { + t.Parallel() + + srv := startUpstream(t) + dialer := &recordingDialer{target: srv.Listener.Addr().String()} + + cfg := DefaultConfig() + cfg.AllowHTTP = true + cfg.DialContext = dialer.dialContext + f := New(cfg) + + if body := fetchBody(t, f, "/image"); body != imagePayload { + t.Errorf("body = %q, want %q", body, imagePayload) + } + + _, err := f.Fetch(testContext(t), "http://127.0.0.1/image") + if !errors.Is(err, ErrSSRFBlocked) { + t.Errorf("Fetch(loopback URL) error = %v, want ErrSSRFBlocked", err) + } + + _, err = f.Fetch(testContext(t), upstreamURL("/redirect/private")) + if !errors.Is(err, ErrSSRFBlocked) { + t.Errorf("Fetch(/redirect/private) error = %v, want ErrSSRFBlocked", err) + } + + // The upstream server is reached through DialContext, and nothing else + // is asked of it. + dialed := dialer.dialedAddrs() + if len(dialed) == 0 { + t.Error("DialContext was never called") + } + + for _, addr := range dialed { + if addr != net.JoinHostPort(testPublicHost, "80") { + t.Errorf("DialContext was asked to connect to %s", addr) + } + } +} diff --git a/internal/httpfetcher/httpfetcher.go b/internal/httpfetcher/httpfetcher.go index f76e8b1..d4ce17f 100644 --- a/internal/httpfetcher/httpfetcher.go +++ b/internal/httpfetcher/httpfetcher.go @@ -137,6 +137,11 @@ type Config struct { // BlockedNetworks are operator-supplied CIDR ranges refused by the // dialer, in addition to the always-enforced built-in ranges. BlockedNetworks []netip.Prefix + // DialContext, when set, makes the fetcher's connections in place of + // the dialer that refuses internal addresses; the URL and redirect + // checks still run. Only tests set it, to reach a local server; the + // config file and the environment cannot. + DialContext func(ctx context.Context, network, addr string) (net.Conn, error) } // DefaultConfig returns a Config with sensible defaults. @@ -193,10 +198,15 @@ func New(config *Config) *HTTPFetcher { // Create transport with SSRF-safe dialer. The dialer re-resolves and // re-checks at connect time (closing the DNS-rebinding window) against // both the built-in ranges and the operator-supplied blocklist. - transport := &http.Transport{ - DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { + dialContext := config.DialContext + if dialContext == nil { + dialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { return dialSSRFSafe(ctx, network, addr, config.BlockedNetworks) - }, + } + } + + transport := &http.Transport{ + DialContext: dialContext, TLSHandshakeTimeout: DefaultTLSTimeout, MaxIdleConns: DefaultMaxIdleConns, IdleConnTimeout: DefaultIdleConnTimeout,