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) } } }