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
+99 -164
View File
@@ -7,6 +7,7 @@ import (
"io" "io"
"net" "net"
"net/http" "net/http"
"net/netip"
"net/url" "net/url"
"slices" "slices"
"strings" "strings"
@@ -32,6 +33,9 @@ const (
tlsHandshakeTimeout = 5 * time.Second tlsHandshakeTimeout = 5 * time.Second
responseHeaderTimeout = 5 * time.Second responseHeaderTimeout = 5 * time.Second
maxContentLength = 10 * 1024 * 1024 maxContentLength = 10 * 1024 * 1024
dnsPinTTL = 5 * time.Minute
maxDNSPinnedHosts = 1024
) )
type TitleFetcher interface { type TitleFetcher interface {
@@ -48,65 +52,66 @@ func (d DefaultDNSResolver) LookupIP(hostname string) ([]net.IP, error) {
return net.LookupIP(hostname) return net.LookupIP(hostname)
} }
type dnsCacheEntry struct {
ips []net.IP
expiresAt time.Time
}
type DNSCache struct { type DNSCache struct {
mu sync.RWMutex mu sync.RWMutex
data map[string][]net.IP data map[string]dnsCacheEntry
} }
func NewDNSCache() *DNSCache { func NewDNSCache() *DNSCache {
return &DNSCache{ return &DNSCache{
data: make(map[string][]net.IP), data: make(map[string]dnsCacheEntry),
} }
} }
func (c *DNSCache) Get(hostname string) ([]net.IP, bool) { func (c *DNSCache) Get(hostname string) ([]net.IP, bool) {
c.mu.RLock() c.mu.RLock()
defer c.mu.RUnlock() defer c.mu.RUnlock()
ips, exists := c.data[hostname] entry, exists := c.data[hostname]
return ips, exists if !exists || time.Now().After(entry.expiresAt) {
return nil, false
}
return entry.ips, true
} }
func (c *DNSCache) Set(hostname string, ips []net.IP) { func (c *DNSCache) Set(hostname string, ips []net.IP) {
c.mu.Lock() c.mu.Lock()
defer c.mu.Unlock() defer c.mu.Unlock()
c.data[hostname] = ips
}
type CachedDNSResolver struct { if len(c.data) >= maxDNSPinnedHosts {
resolver DNSResolver now := time.Now()
cache *DNSCache for host, entry := range c.data {
} if now.After(entry.expiresAt) {
delete(c.data, host)
func NewCachedDNSResolver(resolver DNSResolver) *CachedDNSResolver { }
return &CachedDNSResolver{ }
resolver: resolver, for host := range c.data {
cache: NewDNSCache(), if len(c.data) < maxDNSPinnedHosts {
} break
} }
delete(c.data, host)
func (c *CachedDNSResolver) LookupIP(hostname string) ([]net.IP, error) { }
if ips, exists := c.cache.Get(hostname); exists {
return ips, nil
} }
ips, err := c.resolver.LookupIP(hostname) c.data[hostname] = dnsCacheEntry{
if err != nil { ips: ips,
return nil, err expiresAt: time.Now().Add(dnsPinTTL),
} }
c.cache.Set(hostname, ips)
return ips, nil
} }
type CustomDialer struct { type CustomDialer struct {
cache *DNSCache cache *DNSCache
fallback *net.Dialer dialer *net.Dialer
} }
func NewCustomDialer(cache *DNSCache) *CustomDialer { func NewCustomDialer(cache *DNSCache) *CustomDialer {
return &CustomDialer{ return &CustomDialer{
cache: cache, cache: cache,
fallback: &net.Dialer{ dialer: &net.Dialer{
Timeout: dialTimeout, Timeout: dialTimeout,
}, },
} }
@@ -115,40 +120,43 @@ func NewCustomDialer(cache *DNSCache) *CustomDialer {
func (d *CustomDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { func (d *CustomDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
host, port, err := net.SplitHostPort(address) host, port, err := net.SplitHostPort(address)
if err != nil { if err != nil {
return d.fallback.DialContext(ctx, network, address) return nil, ErrSSRFBlocked
} }
if ips, exists := d.cache.Get(host); exists { ips, exists := d.cache.Get(host)
for _, ip := range ips { if !exists {
ipAddr := net.JoinHostPort(ip.String(), port) return nil, ErrSSRFBlocked
if conn, err := d.fallback.DialContext(ctx, network, ipAddr); err == nil { }
return conn, nil
} var lastErr error
for _, ip := range ips {
conn, err := d.dialer.DialContext(ctx, network, net.JoinHostPort(ip.String(), port))
if err == nil {
return conn, nil
} }
lastErr = err
} }
return d.fallback.DialContext(ctx, network, address) if lastErr == nil {
lastErr = ErrSSRFBlocked
}
return nil, lastErr
} }
type URLMetadataService struct { type URLMetadataService struct {
client *http.Client client *http.Client
resolver DNSResolver resolver DNSResolver
dnsCache *DNSCache dnsCache *DNSCache
approvedHosts map[string]bool
mu sync.RWMutex
} }
func NewURLMetadataService() *URLMetadataService { func NewURLMetadataService() *URLMetadataService {
dnsCache := NewDNSCache()
cachedResolver := NewCachedDNSResolver(DefaultDNSResolver{})
customDialer := NewCustomDialer(dnsCache)
svc := &URLMetadataService{ svc := &URLMetadataService{
resolver: cachedResolver, resolver: DefaultDNSResolver{},
dnsCache: dnsCache, dnsCache: NewDNSCache(),
approvedHosts: make(map[string]bool),
} }
customDialer := NewCustomDialer(svc.dnsCache)
transport := &http.Transport{ transport := &http.Transport{
DialContext: customDialer.DialContext, DialContext: customDialer.DialContext,
MaxIdleConns: 100, MaxIdleConns: 100,
@@ -166,25 +174,7 @@ func NewURLMetadataService() *URLMetadataService {
if len(via) >= maxRedirects { if len(via) >= maxRedirects {
return ErrTooManyRedirects return ErrTooManyRedirects
} }
return svc.validateURLForSSRF(req.URL)
hostname := req.URL.Hostname()
svc.mu.RLock()
approved := svc.approvedHosts[hostname]
svc.mu.RUnlock()
if approved {
return nil
}
if err := svc.validateURLForSSRF(req.URL); err != nil {
return err
}
svc.mu.Lock()
svc.approvedHosts[hostname] = true
svc.mu.Unlock()
return nil
}, },
} }
return svc return svc
@@ -204,19 +194,8 @@ func (s *URLMetadataService) FetchTitle(ctx context.Context, rawURL string) (str
return "", ErrUnsupportedScheme return "", ErrUnsupportedScheme
} }
hostname := parsed.Hostname() if err := s.validateURLForSSRF(parsed); err != nil {
s.mu.RLock() return "", err
approved := s.approvedHosts[hostname]
s.mu.RUnlock()
if !approved {
if err := s.validateURLForSSRF(parsed); err != nil {
return "", err
}
s.mu.Lock()
s.approvedHosts[hostname] = true
s.mu.Unlock()
} }
request, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil) request, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
@@ -336,13 +315,20 @@ func (s *URLMetadataService) validateURLForSSRF(u *url.URL) error {
return ErrSSRFBlocked return ErrSSRFBlocked
} }
ips, err := s.resolver.LookupIP(u.Hostname()) hostname := u.Hostname()
if err != nil { if _, pinned := s.dnsCache.Get(hostname); pinned {
return nil
}
ips, err := s.resolver.LookupIP(hostname)
if err != nil || len(ips) == 0 {
return ErrSSRFBlocked return ErrSSRFBlocked
} }
if slices.ContainsFunc(ips, isPrivateOrReservedIP) { if slices.ContainsFunc(ips, isPrivateOrReservedIP) {
return ErrSSRFBlocked return ErrSSRFBlocked
} }
s.dnsCache.Set(hostname, ips)
return nil return nil
} }
@@ -361,93 +347,42 @@ func isLocalhost(hostname string) bool {
return slices.Contains(localhostNames, hostname) return slices.Contains(localhostNames, hostname)
} }
var reservedPrefixes = []netip.Prefix{
netip.MustParsePrefix("0.0.0.0/8"),
netip.MustParsePrefix("100.64.0.0/10"),
netip.MustParsePrefix("192.0.0.0/24"),
netip.MustParsePrefix("192.0.2.0/24"),
netip.MustParsePrefix("198.18.0.0/15"),
netip.MustParsePrefix("198.51.100.0/24"),
netip.MustParsePrefix("203.0.113.0/24"),
netip.MustParsePrefix("240.0.0.0/4"),
netip.MustParsePrefix("64:ff9b::/96"),
netip.MustParsePrefix("100::/64"),
netip.MustParsePrefix("2001:db8::/32"),
}
func isPrivateOrReservedIP(ip net.IP) bool { func isPrivateOrReservedIP(ip net.IP) bool {
if ip == nil { addr, ok := netip.AddrFromSlice(ip)
if !ok {
return true
}
addr = addr.Unmap()
if addr.IsLoopback() ||
addr.IsPrivate() ||
addr.IsLinkLocalUnicast() ||
addr.IsLinkLocalMulticast() ||
addr.IsInterfaceLocalMulticast() ||
addr.IsMulticast() ||
addr.IsUnspecified() {
return true return true
} }
ipv4 := ip.To4() for _, prefix := range reservedPrefixes {
if ipv4 == nil { if prefix.Contains(addr) {
return isPrivateIPv6(ip)
}
privateRanges := []struct {
start, end net.IP
}{
{net.IPv4(10, 0, 0, 0), net.IPv4(10, 255, 255, 255)},
{net.IPv4(172, 16, 0, 0), net.IPv4(172, 31, 255, 255)},
{net.IPv4(192, 168, 0, 0), net.IPv4(192, 168, 255, 255)},
{net.IPv4(127, 0, 0, 0), net.IPv4(127, 255, 255, 255)},
{net.IPv4(169, 254, 0, 0), net.IPv4(169, 254, 255, 255)},
{net.IPv4(224, 0, 0, 0), net.IPv4(239, 255, 255, 255)},
{net.IPv4(240, 0, 0, 0), net.IPv4(255, 255, 255, 255)},
}
for _, r := range privateRanges {
if ipInRange(ipv4, r.start, r.end) {
return true return true
} }
} }
return false return false
} }
func isPrivateIPv6(ip net.IP) bool {
privateRanges := []struct {
prefix []byte
length int
}{
{[]byte{0xfc, 0x00}, 7},
{[]byte{0xfe, 0x80}, 10},
{[]byte{0xff, 0x00}, 8},
{[]byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01}, 128},
}
for _, r := range privateRanges {
if ipv6InRange(ip, r.prefix, r.length) {
return true
}
}
return false
}
func ipInRange(ip, start, end net.IP) bool {
ipInt := ipToInt(ip)
startInt := ipToInt(start)
endInt := ipToInt(end)
return ipInt >= startInt && ipInt <= endInt
}
func ipToInt(ip net.IP) uint32 {
ipv4 := ip.To4()
if ipv4 == nil {
return 0
}
return uint32(ipv4[0])<<24 + uint32(ipv4[1])<<16 + uint32(ipv4[2])<<8 + uint32(ipv4[3])
}
func ipv6InRange(ip net.IP, prefix []byte, length int) bool {
ipBytes := ip.To16()
if ipBytes == nil {
return false
}
bytesToCompare := length / 8
bitsToCompare := length % 8
for i := 0; i < bytesToCompare && i < len(prefix) && i < len(ipBytes); i++ {
if ipBytes[i] != prefix[i] {
return false
}
}
if bitsToCompare > 0 && bytesToCompare < len(prefix) && bytesToCompare < len(ipBytes) {
mask := byte(0xff) << (8 - bitsToCompare)
if (ipBytes[bytesToCompare] & mask) != (prefix[bytesToCompare] & mask) {
return false
}
}
return true
}
+91 -74
View File
@@ -9,6 +9,7 @@ import (
"net/url" "net/url"
"strings" "strings"
"testing" "testing"
"time"
) )
func TestFetchTitleSuccess(t *testing.T) { func TestFetchTitleSuccess(t *testing.T) {
@@ -428,10 +429,38 @@ func TestIsPrivateOrReservedIP(t *testing.T) {
{"169.254.0.1", "169.254.0.1", true}, {"169.254.0.1", "169.254.0.1", true},
{"224.0.0.1", "224.0.0.1", true}, {"224.0.0.1", "224.0.0.1", true},
{"240.0.0.1", "240.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}, {"8.8.8.8", "8.8.8.8", false},
{"1.1.1.1", "1.1.1.1", false}, {"1.1.1.1", "1.1.1.1", false},
{"74.125.224.72", "74.125.224.72", 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}, {"nil IP", "", true},
} }
@@ -848,93 +877,81 @@ func (c *CountingMockDNSResolver) LookupIP(hostname string) ([]net.IP, error) {
return c.MockDNSResolver.LookupIP(hostname) return c.MockDNSResolver.LookupIP(hostname)
} }
func TestIPv6PrivateRangeDetection(t *testing.T) { func TestCustomDialerRequiresPinnedIP(t *testing.T) {
tests := []struct { cache := NewDNSCache()
name string dialer := NewCustomDialer(cache)
ip string
expected bool _, err := dialer.DialContext(context.Background(), "tcp", "unvalidated.example.com:443")
}{ if !errors.Is(err, ErrSSRFBlocked) {
{"fc00::1", "fc00::1", true}, t.Fatalf("expected ErrSSRFBlocked for unpinned host, got %v", err)
{"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},
} }
for _, tt := range tests { _, err = dialer.DialContext(context.Background(), "tcp", "missing-port.example.com")
t.Run(tt.name, func(t *testing.T) { if !errors.Is(err, ErrSSRFBlocked) {
var ip net.IP t.Fatalf("expected ErrSSRFBlocked for invalid address, got %v", err)
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)
}
})
} }
} }
func TestIPRangeDetection(t *testing.T) { func TestDNSCacheExpiry(t *testing.T) {
tests := []struct { cache := NewDNSCache()
name string cache.Set("example.com", []net.IP{net.ParseIP("8.8.8.8")})
ip string
start string if _, exists := cache.Get("example.com"); !exists {
end string t.Fatal("expected fresh entry to be returned")
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},
} }
for _, tt := range tests { cache.mu.Lock()
t.Run(tt.name, func(t *testing.T) { entry := cache.data["example.com"]
ip := net.ParseIP(tt.ip) entry.expiresAt = time.Now().Add(-time.Second)
start := net.ParseIP(tt.start) cache.data["example.com"] = entry
end := net.ParseIP(tt.end) cache.mu.Unlock()
result := ipInRange(ip, start, end) if _, exists := cache.Get("example.com"); exists {
if result != tt.expected { t.Fatal("expected expired entry to be treated as a miss")
t.Fatalf("expected %v for IP %q in range %q-%q, got %v", tt.expected, tt.ip, tt.start, tt.end, result)
}
})
} }
} }
func TestIPv6RangeDetection(t *testing.T) { func TestExpiredPinTriggersRevalidation(t *testing.T) {
tests := []struct { svc := NewURLMetadataService()
name string
ip string lookupCount := 0
prefix []byte mockResolver := &CountingMockDNSResolver{
length int MockDNSResolver: MockDNSResolver{
expected bool lookupResults: make(map[string][]net.IP),
}{ lookupErrors: make(map[string]error),
{"fc00 prefix match", "fc00::1", []byte{0xfc, 0x00}, 7, true}, },
{"fc00 prefix no match", "fd00::1", []byte{0xfc, 0x00}, 7, true}, lookupCount: &lookupCount,
{"fe80 prefix match", "fe80::1", []byte{0xfe, 0x80}, 10, true}, }
{"fe80 prefix no match", "fe90::1", []byte{0xfe, 0x80}, 10, true}, mockResolver.SetLookupResult("example.com", []net.IP{net.ParseIP("8.8.8.8")})
{"ff00 prefix match", "ff00::1", []byte{0xff, 0x00}, 8, true}, svc.resolver = mockResolver
{"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}, 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 { svc.dnsCache.mu.Lock()
t.Run(tt.name, func(t *testing.T) { entry := svc.dnsCache.data["example.com"]
ip := net.ParseIP(tt.ip) entry.expiresAt = time.Now().Add(-time.Second)
result := ipv6InRange(ip, tt.prefix, tt.length) svc.dnsCache.data["example.com"] = entry
if result != tt.expected { svc.dnsCache.mu.Unlock()
t.Fatalf("expected %v for IPv6 %q with prefix %v/%d, got %v", tt.expected, tt.ip, tt.prefix, tt.length, result)
} 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)
} }
} }