Parent directory

ssrf_test.go

8009 bytes
  1package ssrf
  2
  3import (
  4	"context"
  5	"errors"
  6	"net"
  7	"net/http"
  8	"net/http/httptest"
  9	"testing"
 10	"time"
 11)
 12
 13func TestNewSecureClientBlocksPrivateIPs(t *testing.T) {
 14	client := NewSecureClient()
 15	transport := client.Transport.(*http.Transport)
 16
 17	tests := []struct {
 18		name    string
 19		addr    string
 20		wantErr bool
 21	}{
 22		{"block loopback", "127.0.0.1:80", true},
 23		{"block private 10.x", "10.0.0.1:80", true},
 24		{"block link-local", "169.254.169.254:80", true},
 25		{"block CGNAT", "100.64.0.1:80", true},
 26		{"allow public IP", "8.8.8.8:80", false},
 27	}
 28
 29	for _, tt := range tests {
 30		t.Run(tt.name, func(t *testing.T) {
 31			ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second)
 32			defer cancel()
 33
 34			conn, err := transport.DialContext(ctx, "tcp", tt.addr)
 35			if conn != nil {
 36				conn.Close()
 37			}
 38
 39			if tt.wantErr {
 40				if err == nil {
 41					t.Errorf("expected error for %s, got none", tt.addr)
 42				}
 43			} else {
 44				if err != nil && errors.Is(err, ErrPrivateIP) {
 45					t.Errorf("should not block %s with SSRF error, got: %v", tt.addr, err)
 46				}
 47			}
 48		})
 49	}
 50}
 51
 52func TestNewSecureClientBlocksPrivateHostnames(t *testing.T) {
 53	client := NewSecureClient()
 54	transport := client.Transport.(*http.Transport)
 55
 56	ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
 57	defer cancel()
 58
 59	// "localhost" resolves to 127.0.0.1 (loopback) -- must be blocked.
 60	// This exercises the hostname resolution path in DialContext:
 61	// net.LookupIP("localhost") -> 127.0.0.1 -> isPrivateOrInternal -> blocked.
 62	conn, err := transport.DialContext(ctx, "tcp", "localhost:80")
 63	if conn != nil {
 64		conn.Close()
 65	}
 66	if !errors.Is(err, ErrPrivateIP) {
 67		t.Errorf("expected ErrPrivateIP for hostname resolving to loopback, got: %v", err)
 68	}
 69}
 70
 71func TestNewSecureClientNilIPNotErrPrivateIP(t *testing.T) {
 72	client := NewSecureClient()
 73	transport := client.Transport.(*http.Transport)
 74
 75	ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second)
 76	defer cancel()
 77
 78	conn, err := transport.DialContext(ctx, "tcp", "not-an-ip:80")
 79	if conn != nil {
 80		conn.Close()
 81	}
 82	if err == nil {
 83		t.Fatal("expected error for non-IP address, got none")
 84	}
 85	if errors.Is(err, ErrPrivateIP) {
 86		t.Errorf("nil-IP path should not wrap ErrPrivateIP, got: %v", err)
 87	}
 88}
 89
 90func TestNewSecureClientBlocksRedirects(t *testing.T) {
 91	redirectServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
 92		http.Redirect(w, r, "http://8.8.8.8:8080/safe", http.StatusFound)
 93	}))
 94	defer redirectServer.Close()
 95
 96	client := NewSecureClient()
 97	req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, redirectServer.URL, nil)
 98	if err != nil {
 99		t.Fatalf("Failed to create request: %v", err)
100	}
101
102	resp, err := client.Do(req)
103	if err != nil {
104		// httptest uses 127.0.0.1, blocked by SSRF protection
105		if !errors.Is(err, ErrPrivateIP) {
106			t.Fatalf("Request failed with non-SSRF error: %v", err)
107		}
108		return
109	}
110	defer resp.Body.Close()
111
112	if resp.StatusCode != http.StatusFound {
113		t.Errorf("Expected redirect response (302), got %d", resp.StatusCode)
114	}
115}
116
117func TestIsPrivateOrInternal(t *testing.T) {
118	tests := []struct {
119		ip   string
120		want bool
121	}{
122		// Public
123		{"8.8.8.8", false},
124		{"2001:4860:4860::8888", false},
125
126		// Loopback
127		{"127.0.0.1", true},
128		{"::1", true},
129
130		// Private ranges
131		{"10.0.0.1", true},
132		{"192.168.1.1", true},
133		{"172.16.0.1", true},
134
135		// Link-local (cloud metadata)
136		{"169.254.169.254", true},
137
138		// CGNAT boundaries
139		{"100.64.0.1", true},
140		{"100.127.255.255", true},
141
142		// IPv6-mapped IPv4 (bypass vector the old webhook code missed)
143		{"::ffff:127.0.0.1", true},
144		{"::ffff:169.254.169.254", true},
145		{"::ffff:8.8.8.8", false},
146
147		// Reserved
148		{"0.0.0.0", true},
149		{"240.0.0.1", true},
150
151		// IPv6 transition addresses. Each embeds an IPv4 address inside an
152		// IPv6 one, so they look public to the IPv6 checks while naming an
153		// internal host once a relay decodes them. The whole prefix is
154		// blocked, so the embedded address does not matter: encodings of
155		// public addresses are rejected too.
156		{"2002:7f00:0001::", true},                        // 6to4 -> 127.0.0.1
157		{"2002:a9fe:a9fe::", true},                        // 6to4 -> 169.254.169.254
158		{"2002:0a00:0001::", true},                        // 6to4 -> 10.0.0.1
159		{"2002:0808:0808::", true},                        // 6to4 -> 8.8.8.8, still blocked
160		{"64:ff9b::7f00:1", true},                         // NAT64 well-known -> 127.0.0.1
161		{"64:ff9b::a9fe:a9fe", true},                      // NAT64 well-known -> 169.254.169.254
162		{"64:ff9b:1::7f00:1", true},                       // NAT64 local-use -> 127.0.0.1
163		{"2001:0000:4136:e378:8000:63bf:8001:5601", true}, // Teredo
164		{"::7f00:1", true},                                // IPv4-compatible -> 127.0.0.1
165		{"192.88.99.1", true},                             // 6to4 relay anycast
166
167		// Public IPv6 must stay reachable: the transition prefixes are narrow
168		// and must not swallow ordinary global unicast.
169		{"2606:4700:4700::1111", false},
170		{"2a00:1450:4001::1", false},
171		{"2400:cb00::1", false},
172	}
173
174	for _, tt := range tests {
175		t.Run(tt.ip, func(t *testing.T) {
176			ip := net.ParseIP(tt.ip)
177			if ip == nil {
178				t.Fatalf("failed to parse IP: %s", tt.ip)
179			}
180			if got := isPrivateOrInternal(ip); got != tt.want {
181				t.Errorf("isPrivateOrInternal(%s) = %v, want %v", tt.ip, got, tt.want)
182			}
183		})
184	}
185}
186
187// TestStdlibCoverageParity guards the switch from the standard library's
188// Is*() helpers to an explicit prefix table. The helpers covered these ranges
189// implicitly, so the table has to cover them too or the refactor silently
190// narrows the guard.
191func TestStdlibCoverageParity(t *testing.T) {
192	for _, ip := range []string{
193		// Loopback.
194		"127.0.0.1", "127.255.255.254", "::1",
195		// Private.
196		"10.255.255.255", "172.31.255.255", "192.168.255.255", "fd00::1", "fdff::1",
197		// Link-local unicast.
198		"169.254.1.1", "fe80::1", "febf::1",
199		// Multicast, including link-local and interface-local.
200		"224.0.0.1", "239.255.255.255", "ff00::1", "ff01::1", "ff02::1", "ffff::1",
201		// Unspecified.
202		"0.0.0.0", "::",
203	} {
204		t.Run(ip, func(t *testing.T) {
205			if !isPrivateOrInternal(net.ParseIP(ip)) {
206				t.Errorf("isPrivateOrInternal(%s) = false, want true", ip)
207			}
208		})
209	}
210}
211
212func TestValidateURL(t *testing.T) {
213	tests := []struct {
214		name    string
215		url     string
216		wantErr bool
217		errType error
218	}{
219		// Valid
220		{"valid https", "https://1.1.1.1/webhook", false, nil},
221
222		// Scheme validation
223		{"ftp scheme", "ftp://example.com/webhook", true, ErrInvalidScheme},
224		{"no scheme", "example.com/webhook", true, ErrInvalidScheme},
225
226		// Localhost
227		{"localhost", "http://localhost/webhook", true, ErrPrivateIP},
228		{"subdomain.localhost", "http://test.localhost/webhook", true, ErrPrivateIP},
229
230		// IP-based blocking (one per category -- range coverage is in TestIsPrivateOrInternal)
231		{"loopback IP", "http://127.0.0.1/webhook", true, ErrPrivateIP},
232		{"metadata IP", "http://169.254.169.254/latest/meta-data/", true, ErrPrivateIP},
233
234		// Invalid URLs
235		{"empty", "", true, ErrInvalidURL},
236		{"missing hostname", "http:///webhook", true, ErrInvalidURL},
237	}
238
239	for _, tt := range tests {
240		t.Run(tt.name, func(t *testing.T) {
241			err := ValidateURL(tt.url)
242			if (err != nil) != tt.wantErr {
243				t.Errorf("ValidateURL(%q) error = %v, wantErr %v", tt.url, err, tt.wantErr)
244				return
245			}
246			if tt.wantErr && tt.errType != nil {
247				if !errors.Is(err, tt.errType) {
248					t.Errorf("ValidateURL(%q) error = %v, want error type %v", tt.url, err, tt.errType)
249				}
250			}
251		})
252	}
253}
254
255func TestIsLocalhost(t *testing.T) {
256	tests := []struct {
257		hostname string
258		want     bool
259	}{
260		{"localhost", true},
261		{"LOCALHOST", true},
262		{"test.localhost", true},
263		{"example.com", false},
264		{"localhost.com", false},
265	}
266
267	for _, tt := range tests {
268		t.Run(tt.hostname, func(t *testing.T) {
269			if got := isLocalhost(tt.hostname); got != tt.want {
270				t.Errorf("isLocalhost(%s) = %v, want %v", tt.hostname, got, tt.want)
271			}
272		})
273	}
274}