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}