fix: SSRF DNS rebinding

This commit is contained in:
2026-07-23 16:41:17 +02:00
parent 96958af310
commit e2c46c9e58
2 changed files with 190 additions and 238 deletions
+91 -74
View File
@@ -9,6 +9,7 @@ import (
"net/url"
"strings"
"testing"
"time"
)
func TestFetchTitleSuccess(t *testing.T) {
@@ -428,10 +429,38 @@ func TestIsPrivateOrReservedIP(t *testing.T) {
{"169.254.0.1", "169.254.0.1", true},
{"224.0.0.1", "224.0.0.1", true},
{"240.0.0.1", "240.0.0.1", true},
{"0.0.0.0", "0.0.0.0", true},
{"0.1.2.3", "0.1.2.3", true},
{"CGNAT 100.64.0.1", "100.64.0.1", true},
{"CGNAT 100.127.255.255", "100.127.255.255", true},
{"IETF 192.0.0.1", "192.0.0.1", true},
{"TEST-NET-1 192.0.2.1", "192.0.2.1", true},
{"benchmarking 198.18.0.1", "198.18.0.1", true},
{"benchmarking 198.19.255.255", "198.19.255.255", true},
{"TEST-NET-2 198.51.100.1", "198.51.100.1", true},
{"TEST-NET-3 203.0.113.1", "203.0.113.1", true},
{"broadcast 255.255.255.255", "255.255.255.255", true},
{"IPv4-mapped private ::ffff:10.0.0.1", "::ffff:10.0.0.1", true},
{"IPv4-mapped loopback ::ffff:127.0.0.1", "::ffff:127.0.0.1", true},
{"IPv6 loopback ::1", "::1", true},
{"IPv6 unspecified ::", "::", true},
{"IPv6 ULA fc00::1", "fc00::1", true},
{"IPv6 ULA fd00::1", "fd00::1", true},
{"IPv6 link-local fe80::1", "fe80::1", true},
{"IPv6 multicast ff00::1", "ff00::1", true},
{"IPv6 interface-local multicast ff01::1", "ff01::1", true},
{"NAT64 64:ff9b::7f00:1", "64:ff9b::7f00:1", true},
{"discard-only 100::1", "100::1", true},
{"documentation 2001:db8::1", "2001:db8::1", true},
{"8.8.8.8", "8.8.8.8", false},
{"1.1.1.1", "1.1.1.1", false},
{"74.125.224.72", "74.125.224.72", false},
{"100.128.0.1 just past CGNAT", "100.128.0.1", false},
{"IPv6 public 2001:4860::1", "2001:4860::1", false},
{"IPv6 public 2607:f8b0::1", "2607:f8b0::1", false},
{"IPv4-mapped public ::ffff:8.8.8.8", "::ffff:8.8.8.8", false},
{"nil IP", "", true},
}
@@ -848,93 +877,81 @@ func (c *CountingMockDNSResolver) LookupIP(hostname string) ([]net.IP, error) {
return c.MockDNSResolver.LookupIP(hostname)
}
func TestIPv6PrivateRangeDetection(t *testing.T) {
tests := []struct {
name string
ip string
expected bool
}{
{"fc00::1", "fc00::1", true},
{"fe80::1", "fe80::1", true},
{"ff00::1", "ff00::1", true},
{"::1", "::1", true},
{"2001:db8::1", "2001:db8::1", false},
{"2001:4860::1", "2001:4860::1", false},
{"2607:f8b0::1", "2607:f8b0::1", false},
{"invalid", "invalid", false},
{"", "", false},
func TestCustomDialerRequiresPinnedIP(t *testing.T) {
cache := NewDNSCache()
dialer := NewCustomDialer(cache)
_, err := dialer.DialContext(context.Background(), "tcp", "unvalidated.example.com:443")
if !errors.Is(err, ErrSSRFBlocked) {
t.Fatalf("expected ErrSSRFBlocked for unpinned host, got %v", err)
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var ip net.IP
if tt.ip != "" && tt.ip != "invalid" {
ip = net.ParseIP(tt.ip)
}
result := isPrivateIPv6(ip)
if result != tt.expected {
t.Fatalf("expected %v for IPv6 %q, got %v", tt.expected, tt.ip, result)
}
})
_, err = dialer.DialContext(context.Background(), "tcp", "missing-port.example.com")
if !errors.Is(err, ErrSSRFBlocked) {
t.Fatalf("expected ErrSSRFBlocked for invalid address, got %v", err)
}
}
func TestIPRangeDetection(t *testing.T) {
tests := []struct {
name string
ip string
start string
end string
expected bool
}{
{"IP in range", "192.168.1.100", "192.168.1.1", "192.168.1.255", true},
{"IP at start of range", "192.168.1.1", "192.168.1.1", "192.168.1.255", true},
{"IP at end of range", "192.168.1.255", "192.168.1.1", "192.168.1.255", true},
{"IP below range", "192.168.0.255", "192.168.1.1", "192.168.1.255", false},
{"IP above range", "192.168.2.1", "192.168.1.1", "192.168.1.255", false},
{"Same IP", "192.168.1.100", "192.168.1.100", "192.168.1.100", true},
func TestDNSCacheExpiry(t *testing.T) {
cache := NewDNSCache()
cache.Set("example.com", []net.IP{net.ParseIP("8.8.8.8")})
if _, exists := cache.Get("example.com"); !exists {
t.Fatal("expected fresh entry to be returned")
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ip := net.ParseIP(tt.ip)
start := net.ParseIP(tt.start)
end := net.ParseIP(tt.end)
cache.mu.Lock()
entry := cache.data["example.com"]
entry.expiresAt = time.Now().Add(-time.Second)
cache.data["example.com"] = entry
cache.mu.Unlock()
result := ipInRange(ip, start, end)
if result != tt.expected {
t.Fatalf("expected %v for IP %q in range %q-%q, got %v", tt.expected, tt.ip, tt.start, tt.end, result)
}
})
if _, exists := cache.Get("example.com"); exists {
t.Fatal("expected expired entry to be treated as a miss")
}
}
func TestIPv6RangeDetection(t *testing.T) {
tests := []struct {
name string
ip string
prefix []byte
length int
expected bool
}{
{"fc00 prefix match", "fc00::1", []byte{0xfc, 0x00}, 7, true},
{"fc00 prefix no match", "fd00::1", []byte{0xfc, 0x00}, 7, true},
{"fe80 prefix match", "fe80::1", []byte{0xfe, 0x80}, 10, true},
{"fe80 prefix no match", "fe90::1", []byte{0xfe, 0x80}, 10, true},
{"ff00 prefix match", "ff00::1", []byte{0xff, 0x00}, 8, true},
{"ff00 prefix no match", "fe00::1", []byte{0xff, 0x00}, 8, false},
{"exact match", "::1", []byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01}, 128, true},
func TestExpiredPinTriggersRevalidation(t *testing.T) {
svc := NewURLMetadataService()
lookupCount := 0
mockResolver := &CountingMockDNSResolver{
MockDNSResolver: MockDNSResolver{
lookupResults: make(map[string][]net.IP),
lookupErrors: make(map[string]error),
},
lookupCount: &lookupCount,
}
mockResolver.SetLookupResult("example.com", []net.IP{net.ParseIP("8.8.8.8")})
svc.resolver = mockResolver
svc.client = newTestClient(t, func(r *http.Request) (*http.Response, error) {
body := io.NopCloser(strings.NewReader("<html><head><title>Test Title</title></head></html>"))
header := make(http.Header)
header.Set("Content-Type", "text/html; charset=utf-8")
return &http.Response{StatusCode: http.StatusOK, Body: body, Header: header}, nil
})
if _, err := svc.FetchTitle(context.Background(), "https://example.com"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if lookupCount != 1 {
t.Fatalf("expected 1 DNS lookup, got %d", lookupCount)
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ip := net.ParseIP(tt.ip)
result := ipv6InRange(ip, tt.prefix, tt.length)
if result != tt.expected {
t.Fatalf("expected %v for IPv6 %q with prefix %v/%d, got %v", tt.expected, tt.ip, tt.prefix, tt.length, result)
}
})
svc.dnsCache.mu.Lock()
entry := svc.dnsCache.data["example.com"]
entry.expiresAt = time.Now().Add(-time.Second)
svc.dnsCache.data["example.com"] = entry
svc.dnsCache.mu.Unlock()
mockResolver.SetLookupResult("example.com", []net.IP{net.ParseIP("127.0.0.1")})
if _, err := svc.FetchTitle(context.Background(), "https://example.com"); !errors.Is(err, ErrSSRFBlocked) {
t.Fatalf("expected ErrSSRFBlocked after DNS rebind to private IP, got %v", err)
}
if lookupCount != 2 {
t.Fatalf("expected expired pin to trigger a second DNS lookup, got %d", lookupCount)
}
}