package httpfetcher import ( "context" "errors" "fmt" "io" "net" "net/http" "net/http/httptest" "slices" "strings" "sync" "testing" "time" ) // testPublicHost is a TEST-NET-1 (RFC 5737) literal. isPrivateIP treats it as // public, so validateURL and the redirect check accept it with no DNS lookup, // while the recording dialer routes it to the local httptest server. The // address is reserved for documentation and is never routed on the network. const testPublicHost = "192.0.2.10" // imagePayload is the body served by the fake upstream's image route. const imagePayload = "fake-jpeg-bytes" // errUnexpectedDial reports a dial to any host other than testPublicHost, which // would mean SSRF protection let a forbidden target reach the transport. var errUnexpectedDial = errors.New("unexpected dial target") // upstreamURL builds a fetch URL on the fake public host for the given path. func upstreamURL(path string) string { return "http://" + testPublicHost + path } // recordingDialer records every address the transport asks it to dial and // routes connections for testPublicHost to a real local server, so the SSRF // checks run against a public-looking host while bytes go to httptest. type recordingDialer struct { target string mu sync.Mutex dialed []string } func (d *recordingDialer) dialContext( ctx context.Context, network, addr string, ) (net.Conn, error) { d.mu.Lock() d.dialed = append(d.dialed, addr) d.mu.Unlock() host, _, err := net.SplitHostPort(addr) if err != nil { return nil, err } if host != testPublicHost { return nil, fmt.Errorf("%w: %s", errUnexpectedDial, addr) } var dialer net.Dialer return dialer.DialContext(ctx, network, d.target) } // dialedAddrs returns a copy of the addresses the dialer was asked to reach. func (d *recordingDialer) dialedAddrs() []string { d.mu.Lock() defer d.mu.Unlock() return slices.Clone(d.dialed) } // startUpstream launches a fake upstream with the routes the fetch tests // exercise and stops it when the test finishes. func startUpstream(t *testing.T) *httptest.Server { t.Helper() mux := http.NewServeMux() mux.HandleFunc("/image", func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", contentTypeJPEG) _, _ = io.WriteString(w, imagePayload) }) mux.HandleFunc("/status/500", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusInternalServerError) }) mux.HandleFunc("/html", func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "text/html; charset=utf-8") _, _ = io.WriteString(w, "") }) mux.HandleFunc("/redirect/private", func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "http://169.254.169.254/latest/meta-data/", http.StatusFound) }) mux.HandleFunc("/redirect/public", func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/image", http.StatusFound) }) mux.HandleFunc("/redirect/chain", func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/redirect/hop", http.StatusFound) }) mux.HandleFunc("/redirect/hop", func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/image", http.StatusFound) }) srv := httptest.NewServer(mux) t.Cleanup(srv.Close) return srv } // newServerFetcher builds a fetcher whose transport routes testPublicHost to // srv, leaving the real SSRF validation and redirect checks in place. func newServerFetcher( t *testing.T, srv *httptest.Server, cfg *Config, ) (*HTTPFetcher, *recordingDialer) { t.Helper() if cfg == nil { cfg = DefaultConfig() } cfg.AllowHTTP = true f := New(cfg) transport, ok := f.client.Transport.(*http.Transport) if !ok { t.Fatalf("transport is %T, want *http.Transport", f.client.Transport) } dialer := &recordingDialer{target: srv.Listener.Addr().String()} transport.DialContext = dialer.dialContext return f, dialer } // testContext returns a context cancelled when the test ends, bounding any // fetch that would otherwise block on a leaked semaphore slot. func testContext(t *testing.T) context.Context { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) t.Cleanup(cancel) return ctx } // fetchImage fetches path from the fake upstream and fails on error. func fetchImage(t *testing.T, f *HTTPFetcher, path string) *FetchResult { t.Helper() res, err := f.Fetch(testContext(t), upstreamURL(path)) if err != nil { t.Fatalf("Fetch(%s) error = %v", path, err) } return res } // fetchExpectError fetches path and fails unless Fetch returns an error. func fetchExpectError(t *testing.T, f *HTTPFetcher, path string) error { t.Helper() res, err := f.Fetch(testContext(t), upstreamURL(path)) if err == nil { _ = res.Content.Close() t.Fatalf("Fetch(%s) = nil error, want an error", path) } return err } // fetchBody fetches path and returns the fully read, closed response body. func fetchBody(t *testing.T, f *HTTPFetcher, path string) string { t.Helper() res := fetchImage(t, f, path) defer func() { _ = res.Content.Close() }() data, err := io.ReadAll(res.Content) if err != nil { t.Fatalf("read body: %v", err) } return string(data) } // semLen reports how many per-host semaphore slots are currently held. func semLen(f *HTTPFetcher, host string) int { return len(f.getHostSemaphore(host)) } func TestFetchRedirectToPrivateIPBlocked(t *testing.T) { t.Parallel() srv := startUpstream(t) f, dialer := newServerFetcher(t, srv, nil) _, err := f.Fetch(testContext(t), upstreamURL("/redirect/private")) if !errors.Is(err, ErrSSRFBlocked) { t.Fatalf("Fetch() error = %v, want ErrSSRFBlocked", err) } for _, addr := range dialer.dialedAddrs() { if strings.Contains(addr, "169.254.169.254") { t.Errorf("dialer connected to the private redirect target: %s", addr) } } } func TestFetchRedirectToPublicSucceeds(t *testing.T) { t.Parallel() srv := startUpstream(t) f, _ := newServerFetcher(t, srv, nil) if body := fetchBody(t, f, "/redirect/public"); body != imagePayload { t.Errorf("body = %q, want %q", body, imagePayload) } } func TestFetchRedirectChainSucceeds(t *testing.T) { t.Parallel() srv := startUpstream(t) f, _ := newServerFetcher(t, srv, nil) if body := fetchBody(t, f, "/redirect/chain"); body != imagePayload { t.Errorf("body = %q, want %q", body, imagePayload) } } func TestFetchRejectsNon2xx(t *testing.T) { t.Parallel() srv := startUpstream(t) f, _ := newServerFetcher(t, srv, nil) err := fetchExpectError(t, f, "/status/500") if !errors.Is(err, ErrUpstreamError) { t.Fatalf("Fetch() error = %v, want ErrUpstreamError", err) } } func TestFetchRejectsDisallowedContentType(t *testing.T) { t.Parallel() srv := startUpstream(t) f, _ := newServerFetcher(t, srv, nil) err := fetchExpectError(t, f, "/html") if !errors.Is(err, ErrInvalidContentType) { t.Fatalf("Fetch() error = %v, want ErrInvalidContentType", err) } } func TestFetchMaxResponseSizeEnforced(t *testing.T) { t.Parallel() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxResponseSize = 8 f, _ := newServerFetcher(t, srv, cfg) res := fetchImage(t, f, "/image") defer func() { _ = res.Content.Close() }() data, err := io.ReadAll(res.Content) if !errors.Is(err, ErrResponseTooLarge) { t.Fatalf("read error = %v, want ErrResponseTooLarge", err) } if int64(len(data)) > cfg.MaxResponseSize { t.Errorf("read %d bytes, exceeds limit %d", len(data), cfg.MaxResponseSize) } } func TestFetchSemaphoreReleasedOnError(t *testing.T) { t.Parallel() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxConnectionsPerHost = 1 f, _ := newServerFetcher(t, srv, cfg) err := fetchExpectError(t, f, "/status/500") if !errors.Is(err, ErrUpstreamError) { t.Fatalf("Fetch() error = %v, want ErrUpstreamError", err) } if held := semLen(f, testPublicHost); held != 0 { t.Fatalf("semaphore slot leaked after error: %d held", held) } // One slot per host: this fetch proceeds only if the slot was released. res := fetchImage(t, f, "/image") _ = res.Content.Close() } // assertSlotReleasedByClose fetches an image over a one-slot host, hands the // open result to consume, and asserts the slot is held before and freed after, // then that a follow-up fetch can still acquire it. func assertSlotReleasedByClose( t *testing.T, consume func(*testing.T, *FetchResult), ) { t.Helper() srv := startUpstream(t) cfg := DefaultConfig() cfg.MaxConnectionsPerHost = 1 f, _ := newServerFetcher(t, srv, cfg) res := fetchImage(t, f, "/image") if held := semLen(f, testPublicHost); held != 1 { t.Fatalf("slot not held while body is open: %d held", held) } consume(t, res) if held := semLen(f, testPublicHost); held != 0 { t.Fatalf("slot not released after close: %d held", held) } next := fetchImage(t, f, "/image") _ = next.Content.Close() } func TestFetchSemaphoreReleasedOnBodyClose(t *testing.T) { t.Parallel() assertSlotReleasedByClose(t, func(t *testing.T, res *FetchResult) { t.Helper() _, err := io.ReadAll(res.Content) if err != nil { t.Fatalf("read body: %v", err) } err = res.Content.Close() if err != nil { t.Fatalf("close body: %v", err) } }) } func TestFetchSemaphoreReleasedOnPartialReadClose(t *testing.T) { t.Parallel() assertSlotReleasedByClose(t, func(t *testing.T, res *FetchResult) { t.Helper() buf := make([]byte, 1) _, err := res.Content.Read(buf) if err != nil { t.Fatalf("partial read: %v", err) } err = res.Content.Close() if err != nil { t.Fatalf("close body: %v", err) } }) } // The dial-time re-resolution in ssrfSafeDialer is what closes the DNS // rebinding window: even if validateURL saw a public answer earlier, the // dialer independently re-checks the address it is about to connect to. A full // rebinding simulation (a resolver returning public, then private) would mean // replacing the global net.DefaultResolver with a fake DNS server, which is // heavyweight and unsafe to mutate under parallel -race tests. The property is // proven directly here instead: the dialer rejects a private target outright, // which is exactly the check that fires when a validated host later resolves // to a private address. func TestSSRFSafeDialerBlocksPrivateTarget(t *testing.T) { t.Parallel() for _, addr := range []string{ "169.254.169.254:80", // link-local (cloud metadata) "127.0.0.1:80", // loopback "10.0.0.5:80", // RFC 1918 private } { t.Run(addr, func(t *testing.T) { t.Parallel() _, err := ssrfSafeDialer(context.Background(), "tcp", addr) if !errors.Is(err, ErrSSRFBlocked) { t.Errorf("ssrfSafeDialer(%q) = %v, want ErrSSRFBlocked", addr, err) } }) } } func TestSSRFSafeDialerAllowsPublicTarget(t *testing.T) { t.Parallel() // A cancelled context makes the dial fail immediately without touching the // network; the point is only that a public literal is not SSRF-blocked. ctx, cancel := context.WithCancel(context.Background()) cancel() _, err := ssrfSafeDialer(ctx, "tcp", testPublicHost+":80") if err == nil { t.Fatal("expected a dial error for an unreachable public target") } if errors.Is(err, ErrSSRFBlocked) { t.Errorf("public target was SSRF-blocked: %v", err) } }