From 3b1d7be6fcd8d6c75dc368fb45e4e22aface5dfe Mon Sep 17 00:00:00 2001 From: Dmitry Verkhoturov Date: Sat, 18 Apr 2026 05:07:11 +0100 Subject: [PATCH] fix(safehttp): clone http.DefaultTransport, sharpen Image.Transport contract Address review feedback on PR #2044. safehttp.Transport(): * Clone http.DefaultTransport instead of building a bare &http.Transport{} so Proxy, ForceAttemptHTTP2, MaxIdleConns, IdleConnTimeout, TLSHandshakeTimeout and ExpectContinueTimeout are inherited (the bare struct loses them all). Verified by new TestTransport_PreservesDefaultTransportSettings. * TestTransport_AllowsPublic: bound the dial of TEST-NET-3 with a 100ms context so the test does not depend on real-world routing of 203.0.113.0/24, and drop the dead dialer var. proxy/image.go: * Document Image.Transport contract: nil installs safehttp.Transport (SSRF-safe); caller-supplied transport is the caller's responsibility. * Replace the misleading "SSRF mitigated by safehttp.Transport" nolint comments with one that points at the documented contract above. --- backend/app/rest/proxy/image.go | 13 +++++-- backend/app/safehttp/safehttp.go | 56 ++++++++++++++------------- backend/app/safehttp/safehttp_test.go | 22 ++++++++--- 3 files changed, 57 insertions(+), 34 deletions(-) diff --git a/backend/app/rest/proxy/image.go b/backend/app/rest/proxy/image.go index 2589e099..3dbb5a8b 100644 --- a/backend/app/rest/proxy/image.go +++ b/backend/app/rest/proxy/image.go @@ -28,7 +28,11 @@ type Image struct { CacheExternal bool Timeout time.Duration ImageService *image.Service - Transport http.RoundTripper // if nil, uses SSRF-safe transport blocking private IPs + // Transport, if non-nil, is used as-is for outbound image fetches and is the + // caller's responsibility to make SSRF-safe. When nil, safehttp.Transport() + // is installed, which blocks dialing any private/reserved IP and resolves + // hostnames to defeat DNS rebinding. + Transport http.RoundTripper } // Convert img src links to proxied links depends on enabled options @@ -163,11 +167,14 @@ func (p Image) downloadImage(ctx context.Context, imgURL string) ([]byte, error) var resp *http.Response err := repeater.NewFixed(5, time.Second).Do(ctx, func() error { var e error - req, e := http.NewRequest("GET", imgURL, http.NoBody) //nolint:gosec // SSRF mitigated by safehttp.Transport assigned above + // SSRF safety: client.Transport is safehttp.Transport() when p.Transport is nil + // (see Image.Transport contract above); when caller supplies a transport they + // own SSRF safety for that path. + req, e := http.NewRequest("GET", imgURL, http.NoBody) //nolint:gosec // see comment above if e != nil { return fmt.Errorf("failed to make request for %s: %w", imgURL, e) } - resp, e = client.Do(req.WithContext(ctx)) //nolint:bodyclose,gosec // body closed in defer; SSRF mitigated by safehttp.Transport + resp, e = client.Do(req.WithContext(ctx)) //nolint:bodyclose,gosec // body closed in defer; transport contract above return e }) if err != nil { diff --git a/backend/app/safehttp/safehttp.go b/backend/app/safehttp/safehttp.go index faf4adc5..a60a69ae 100644 --- a/backend/app/safehttp/safehttp.go +++ b/backend/app/safehttp/safehttp.go @@ -16,40 +16,44 @@ import ( // that resolves to a private/reserved IP, choosing the IP itself for the dial // to defeat DNS rebinding (an attacker cannot have the resolver hand back a // public IP at the check and a private one at the connect). +// +// The returned transport is a clone of http.DefaultTransport with only +// DialContext overridden, preserving Proxy, HTTP/2, idle/keep-alive and +// TLS handshake timeouts that bare &http.Transport{} would lose. func Transport() *http.Transport { dialer := &net.Dialer{Timeout: 30 * time.Second} - return &http.Transport{ - DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { - host, port, err := net.SplitHostPort(addr) - if err != nil { - return nil, fmt.Errorf("invalid address %s: %w", addr, err) - } + t := http.DefaultTransport.(*http.Transport).Clone() + t.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + host, port, err := net.SplitHostPort(addr) + if err != nil { + return nil, fmt.Errorf("invalid address %s: %w", addr, err) + } - ips, err := net.DefaultResolver.LookupIPAddr(ctx, host) - if err != nil { - return nil, fmt.Errorf("can't resolve host %s: %w", host, err) - } - if len(ips) == 0 { - return nil, fmt.Errorf("no IP addresses resolved for host %s", host) - } + ips, err := net.DefaultResolver.LookupIPAddr(ctx, host) + if err != nil { + return nil, fmt.Errorf("can't resolve host %s: %w", host, err) + } + if len(ips) == 0 { + return nil, fmt.Errorf("no IP addresses resolved for host %s", host) + } - for _, ip := range ips { - if IsPrivateIP(ip.IP) { - return nil, fmt.Errorf("access to private address is not allowed") - } + for _, ip := range ips { + if IsPrivateIP(ip.IP) { + return nil, fmt.Errorf("access to private address is not allowed") } + } - var lastErr error - for _, ip := range ips { - conn, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port)) - if dialErr == nil { - return conn, nil - } - lastErr = dialErr + var lastErr error + for _, ip := range ips { + conn, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port)) + if dialErr == nil { + return conn, nil } - return nil, fmt.Errorf("can't connect to %s: %w", host, lastErr) - }, + lastErr = dialErr + } + return nil, fmt.Errorf("can't connect to %s: %w", host, lastErr) } + return t } // privateCIDRs holds pre-parsed private/reserved CIDR blocks. diff --git a/backend/app/safehttp/safehttp_test.go b/backend/app/safehttp/safehttp_test.go index b3fab00e..35a6325a 100644 --- a/backend/app/safehttp/safehttp_test.go +++ b/backend/app/safehttp/safehttp_test.go @@ -63,16 +63,28 @@ func TestTransport_BlocksPrivate(t *testing.T) { } func TestTransport_AllowsPublic(t *testing.T) { - dialer := &net.Dialer{Timeout: 2 * time.Second} tr := Transport() - // monkey-check the Dialer wiring directly: the transport must reject 127.0.0.1 + // the policy check must reject the loopback literal _, err := tr.DialContext(context.Background(), "tcp", "127.0.0.1:1") require.Error(t, err) assert.Contains(t, err.Error(), "access to private address is not allowed") - // public IP literal goes through the dial path (will likely error on connect, but NOT on policy) - _, err = tr.DialContext(context.Background(), "tcp", "203.0.113.1:1") + // public IP literal passes the policy check; bound the dial with a tight context + // so the test does not depend on real-world routing of TEST-NET-3 (203.0.113.0/24). + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + _, err = tr.DialContext(ctx, "tcp", "203.0.113.1:1") require.Error(t, err) assert.NotContains(t, err.Error(), "access to private address is not allowed") - _ = dialer +} + +func TestTransport_PreservesDefaultTransportSettings(t *testing.T) { + def := http.DefaultTransport.(*http.Transport) + tr := Transport() + assert.NotNil(t, tr.Proxy, "Proxy must be inherited from http.DefaultTransport") + assert.Equal(t, def.ForceAttemptHTTP2, tr.ForceAttemptHTTP2, "ForceAttemptHTTP2") + assert.Equal(t, def.MaxIdleConns, tr.MaxIdleConns, "MaxIdleConns") + assert.Equal(t, def.IdleConnTimeout, tr.IdleConnTimeout, "IdleConnTimeout") + assert.Equal(t, def.TLSHandshakeTimeout, tr.TLSHandshakeTimeout, "TLSHandshakeTimeout") + assert.Equal(t, def.ExpectContinueTimeout, tr.ExpectContinueTimeout, "ExpectContinueTimeout") }