Parent directory

ssrf.go

7274 bytes
  1package ssrf
  2
  3import (
  4	"context"
  5	"errors"
  6	"fmt"
  7	"net"
  8	"net/http"
  9	"net/netip"
 10	"net/url"
 11	"slices"
 12	"strings"
 13	"time"
 14)
 15
 16var (
 17	// ErrPrivateIP is returned when a connection to a private or internal IP is blocked.
 18	ErrPrivateIP = errors.New("connection to private or internal IP address is not allowed")
 19	// ErrInvalidScheme is returned when a URL scheme is not http or https.
 20	ErrInvalidScheme = errors.New("URL must use http or https scheme")
 21	// ErrInvalidURL is returned when a URL is invalid.
 22	ErrInvalidURL = errors.New("invalid URL")
 23)
 24
 25// NewSecureClient returns an HTTP client with SSRF protection.
 26// It validates resolved IPs at dial time to block connections to private
 27// and internal networks. Hostnames are resolved and the validated IP is
 28// used directly in the dial call to prevent DNS rebinding (TOCTOU between
 29// validation and connection). Redirects are disabled to match the webhook
 30// client convention and prevent redirect-based SSRF.
 31func NewSecureClient() *http.Client {
 32	return &http.Client{
 33		Timeout: 30 * time.Second,
 34		Transport: &http.Transport{
 35			DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
 36				host, port, err := net.SplitHostPort(addr)
 37				if err != nil {
 38					return nil, err //nolint:wrapcheck
 39				}
 40
 41				ip := net.ParseIP(host)
 42				if ip == nil {
 43					ips, err := net.LookupIP(host) //nolint
 44					if err != nil {
 45						return nil, fmt.Errorf("DNS resolution failed for host %s: %v", host, err)
 46					}
 47					if len(ips) == 0 {
 48						return nil, fmt.Errorf("no IP addresses found for host: %s", host)
 49					}
 50					ip = ips[0] // Use the first resolved IP address
 51				}
 52				if isPrivateOrInternal(ip) {
 53					return nil, fmt.Errorf("%w", ErrPrivateIP)
 54				}
 55
 56				dialer := &net.Dialer{
 57					Timeout:   10 * time.Second,
 58					KeepAlive: 30 * time.Second,
 59				}
 60				// Dial using the validated IP to prevent DNS rebinding.
 61				// Without this, the dialer resolves the hostname again
 62				// independently, and the second resolution could return
 63				// a different (private) IP.
 64				return dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port))
 65			},
 66			MaxIdleConns:          100,
 67			IdleConnTimeout:       90 * time.Second,
 68			TLSHandshakeTimeout:   10 * time.Second,
 69			ExpectContinueTimeout: 1 * time.Second,
 70		},
 71		CheckRedirect: func(*http.Request, []*http.Request) error {
 72			return http.ErrUseLastResponse
 73		},
 74	}
 75}
 76
 77// blockedPrefixes lists every network an outbound request must not reach.
 78//
 79// Prefixes are used rather than byte comparisons so the ranges stay readable
 80// and auditable against the RFCs that define them. This is also how the
 81// standard library models the same address space internally.
 82//
 83// The IPv6 transition ranges need explanation. Each embeds an arbitrary IPv4
 84// address inside an IPv6 one, so an address that looks public to the IPv6
 85// checks can name a loopback, RFC1918, or cloud metadata host once a relay
 86// decodes it. They are blocked in full rather than decoded and re-checked,
 87// because nothing here has a legitimate reason to reach a host through one of
 88// these encodings. Blocking the range removes the whole class of bypass
 89// instead of depending on getting each decoding exactly right.
 90var blockedPrefixes = []netip.Prefix{
 91	// IPv4.
 92	netip.MustParsePrefix("0.0.0.0/8"),       // "this network"
 93	netip.MustParsePrefix("10.0.0.0/8"),      // RFC1918 private
 94	netip.MustParsePrefix("100.64.0.0/10"),   // RFC6598 shared address space (CGNAT)
 95	netip.MustParsePrefix("127.0.0.0/8"),     // loopback
 96	netip.MustParsePrefix("169.254.0.0/16"),  // link-local, includes cloud metadata
 97	netip.MustParsePrefix("172.16.0.0/12"),   // RFC1918 private
 98	netip.MustParsePrefix("192.0.0.0/24"),    // IETF protocol assignments
 99	netip.MustParsePrefix("192.0.2.0/24"),    // TEST-NET-1
100	netip.MustParsePrefix("192.88.99.0/24"),  // RFC7526 6to4 relay anycast
101	netip.MustParsePrefix("192.168.0.0/16"),  // RFC1918 private
102	netip.MustParsePrefix("198.18.0.0/15"),   // RFC2544 benchmarking
103	netip.MustParsePrefix("198.51.100.0/24"), // TEST-NET-2
104	netip.MustParsePrefix("203.0.113.0/24"),  // TEST-NET-3
105	netip.MustParsePrefix("224.0.0.0/4"),     // multicast
106	netip.MustParsePrefix("240.0.0.0/4"),     // reserved, includes broadcast
107
108	// IPv6.
109	netip.MustParsePrefix("::1/128"),        // loopback
110	netip.MustParsePrefix("64:ff9b::/96"),   // RFC6052 NAT64 well-known prefix
111	netip.MustParsePrefix("64:ff9b:1::/48"), // RFC8215 NAT64 local-use prefix
112	netip.MustParsePrefix("100::/64"),       // RFC6666 discard-only
113	netip.MustParsePrefix("2001::/32"),      // RFC4380 Teredo
114	netip.MustParsePrefix("2001:db8::/32"),  // documentation
115	netip.MustParsePrefix("2002::/16"),      // RFC3056 6to4
116	netip.MustParsePrefix("fc00::/7"),       // unique local
117	netip.MustParsePrefix("fe80::/10"),      // link-local
118	netip.MustParsePrefix("ff00::/8"),       // multicast
119
120	// IPv4-compatible IPv6 (::x.x.x.x), deprecated by RFC4291 but still
121	// parsed. Unmap leaves this form alone, so the IPv4 prefixes above do not
122	// apply to it. Covers the unspecified address too.
123	netip.MustParsePrefix("::/96"),
124}
125
126// isPrivateOrInternal reports whether an IP address is private, internal, or
127// otherwise unsafe to send an outbound request to.
128func isPrivateOrInternal(ip net.IP) bool {
129	addr, ok := netip.AddrFromSlice(ip)
130	if !ok {
131		// Not an address we can reason about, so refuse rather than allow.
132		return true
133	}
134
135	// Treat an IPv4 address written in IPv6 form (::ffff:127.0.0.1) as the
136	// IPv4 address it names, so the IPv4 prefixes apply to it.
137	addr = addr.Unmap()
138
139	return slices.ContainsFunc(blockedPrefixes, func(p netip.Prefix) bool {
140		return p.Contains(addr)
141	})
142}
143
144// ValidateURL validates that a URL is safe to make requests to.
145// It checks that the scheme is http/https, the hostname is not localhost,
146// and all resolved IPs are public.
147func ValidateURL(rawURL string) error {
148	if rawURL == "" {
149		return ErrInvalidURL
150	}
151
152	u, err := url.Parse(rawURL)
153	if err != nil {
154		return fmt.Errorf("%w: %v", ErrInvalidURL, err)
155	}
156
157	if u.Scheme != "http" && u.Scheme != "https" {
158		return ErrInvalidScheme
159	}
160
161	hostname := u.Hostname()
162	if hostname == "" {
163		return fmt.Errorf("%w: missing hostname", ErrInvalidURL)
164	}
165
166	if isLocalhost(hostname) {
167		return ErrPrivateIP
168	}
169
170	if ip := net.ParseIP(hostname); ip != nil {
171		if isPrivateOrInternal(ip) {
172			return ErrPrivateIP
173		}
174		return nil
175	}
176
177	ips, err := net.DefaultResolver.LookupIPAddr(context.Background(), hostname)
178	if err != nil {
179		return fmt.Errorf("%w: cannot resolve hostname: %v", ErrInvalidURL, err)
180	}
181
182	if slices.ContainsFunc(ips, func(addr net.IPAddr) bool {
183		return isPrivateOrInternal(addr.IP)
184	}) {
185		return ErrPrivateIP
186	}
187
188	return nil
189}
190
191// ValidateIPBeforeDial validates an IP address before establishing a connection.
192// This prevents DNS rebinding attacks by checking the resolved IP at dial time.
193func ValidateIPBeforeDial(ip net.IP) error {
194	if isPrivateOrInternal(ip) {
195		return ErrPrivateIP
196	}
197	return nil
198}
199
200// isLocalhost checks if the hostname is localhost or similar.
201func isLocalhost(hostname string) bool {
202	hostname = strings.ToLower(hostname)
203	return hostname == "localhost" ||
204		hostname == "localhost.localdomain" ||
205		strings.HasSuffix(hostname, ".localhost")
206}