From eb75b92fb078f0c0c20545dc6fb6c4264d188458 Mon Sep 17 00:00:00 2001 From: sneak Date: Mon, 21 Sep 2026 07:58:26 +0000 Subject: [PATCH] test: cover redirect SSRF and semaphore release in httpfetcher (closes #78) Adds httptest-server tests for the risk-bearing paths that only had helper-level coverage: the CheckRedirect validator (a 302 to a link-local address is refused with ErrSSRFBlocked and never dialed, while a redirect chain and a redirect to a public target still succeed), per-host semaphore release on the error, full-read-close, and partial-read-close paths (proven by saturating a one-slot host), the MaxResponseSize limit end-to-end through Fetch, non-2xx and disallowed content-type rejection, and ssrfSafeDialer's dial-time block of a private target. To reach a loopback test server while the real SSRF checks run, the upstream host is a TEST-NET-1 literal (192.0.2.10) that isPrivateIP treats as public and that resolves with no DNS, and a recording dialer routes it to the server. Full DNS-rebinding simulation is documented as out of scope; the dial-time private-target block that closes that window is tested directly. model: claude-opus-4-8 --- internal/httpfetcher/fetch_internal_test.go | 421 ++++++++++++++++++++ 1 file changed, 421 insertions(+) create mode 100644 internal/httpfetcher/fetch_internal_test.go diff --git a/internal/httpfetcher/fetch_internal_test.go b/internal/httpfetcher/fetch_internal_test.go new file mode 100644 index 0000000..c0d6e9d --- /dev/null +++ b/internal/httpfetcher/fetch_internal_test.go @@ -0,0 +1,421 @@ +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) + } +}