879ece7c24bc92b25e1fb0a4da33049ed8b8332a

Author
Vinayak Mishra <viks@vnykmshr.com>
Committer
GitHub <noreply@github.com>
Date

Message

fix(ssrf): pin resolved IP in dial to prevent DNS rebinding (#791)

Diff

 1diff --git a/pkg/ssrf/ssrf.go b/pkg/ssrf/ssrf.go
 2index 6a94d72e564c98698d8c43330f005b4ce253ec42..475c1cd7216eb7400a01dc39227b8ac818af64cb 100644
 3--- a/pkg/ssrf/ssrf.go
 4+++ b/pkg/ssrf/ssrf.go
 5@@ -23,16 +23,16 @@ var (
 6 
 7 // NewSecureClient returns an HTTP client with SSRF protection.
 8 // It validates resolved IPs at dial time to block connections to private
 9-// and internal networks. Since validation uses the already-resolved IP
10-// from the Transport's DNS lookup, there is no TOCTOU gap between
11-// resolution and connection. Redirects are disabled to match the
12-// webhook client convention and prevent redirect-based SSRF.
13+// and internal networks. Hostnames are resolved and the validated IP is
14+// used directly in the dial call to prevent DNS rebinding (TOCTOU between
15+// validation and connection). Redirects are disabled to match the webhook
16+// client convention and prevent redirect-based SSRF.
17 func NewSecureClient() *http.Client {
18 	return &http.Client{
19 		Timeout: 30 * time.Second,
20 		Transport: &http.Transport{
21 			DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
22-				host, _, err := net.SplitHostPort(addr)
23+				host, port, err := net.SplitHostPort(addr)
24 				if err != nil {
25 					return nil, err //nolint:wrapcheck
26 				}
27@@ -56,7 +56,11 @@ func NewSecureClient() *http.Client {
28 					Timeout:   10 * time.Second,
29 					KeepAlive: 30 * time.Second,
30 				}
31-				return dialer.DialContext(ctx, network, addr)
32+				// Dial using the validated IP to prevent DNS rebinding.
33+				// Without this, the dialer resolves the hostname again
34+				// independently, and the second resolution could return
35+				// a different (private) IP.
36+				return dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port))
37 			},
38 			MaxIdleConns:          100,
39 			IdleConnTimeout:       90 * time.Second,
40diff --git a/pkg/ssrf/ssrf_test.go b/pkg/ssrf/ssrf_test.go
41index a3c684dcf1babf37855c20a734aceeda6ceb1107..1b14194695a6d694221dd2a15050dd8027349720 100644
42--- a/pkg/ssrf/ssrf_test.go
43+++ b/pkg/ssrf/ssrf_test.go
44@@ -49,6 +49,25 @@ func TestNewSecureClientBlocksPrivateIPs(t *testing.T) {
45 	}
46 }
47 
48+func TestNewSecureClientBlocksPrivateHostnames(t *testing.T) {
49+	client := NewSecureClient()
50+	transport := client.Transport.(*http.Transport)
51+
52+	ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
53+	defer cancel()
54+
55+	// "localhost" resolves to 127.0.0.1 (loopback) -- must be blocked.
56+	// This exercises the hostname resolution path in DialContext:
57+	// net.LookupIP("localhost") -> 127.0.0.1 -> isPrivateOrInternal -> blocked.
58+	conn, err := transport.DialContext(ctx, "tcp", "localhost:80")
59+	if conn != nil {
60+		conn.Close()
61+	}
62+	if !errors.Is(err, ErrPrivateIP) {
63+		t.Errorf("expected ErrPrivateIP for hostname resolving to loopback, got: %v", err)
64+	}
65+}
66+
67 func TestNewSecureClientNilIPNotErrPrivateIP(t *testing.T) {
68 	client := NewSecureClient()
69 	transport := client.Transport.(*http.Transport)