fix: SSRF DNS rebinding
This commit is contained in:
@@ -7,6 +7,7 @@ import (
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strings"
|
||||
@@ -32,6 +33,9 @@ const (
|
||||
tlsHandshakeTimeout = 5 * time.Second
|
||||
responseHeaderTimeout = 5 * time.Second
|
||||
maxContentLength = 10 * 1024 * 1024
|
||||
|
||||
dnsPinTTL = 5 * time.Minute
|
||||
maxDNSPinnedHosts = 1024
|
||||
)
|
||||
|
||||
type TitleFetcher interface {
|
||||
@@ -48,65 +52,66 @@ func (d DefaultDNSResolver) LookupIP(hostname string) ([]net.IP, error) {
|
||||
return net.LookupIP(hostname)
|
||||
}
|
||||
|
||||
type dnsCacheEntry struct {
|
||||
ips []net.IP
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type DNSCache struct {
|
||||
mu sync.RWMutex
|
||||
data map[string][]net.IP
|
||||
data map[string]dnsCacheEntry
|
||||
}
|
||||
|
||||
func NewDNSCache() *DNSCache {
|
||||
return &DNSCache{
|
||||
data: make(map[string][]net.IP),
|
||||
data: make(map[string]dnsCacheEntry),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *DNSCache) Get(hostname string) ([]net.IP, bool) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
ips, exists := c.data[hostname]
|
||||
return ips, exists
|
||||
entry, exists := c.data[hostname]
|
||||
if !exists || time.Now().After(entry.expiresAt) {
|
||||
return nil, false
|
||||
}
|
||||
return entry.ips, true
|
||||
}
|
||||
|
||||
func (c *DNSCache) Set(hostname string, ips []net.IP) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.data[hostname] = ips
|
||||
}
|
||||
|
||||
type CachedDNSResolver struct {
|
||||
resolver DNSResolver
|
||||
cache *DNSCache
|
||||
}
|
||||
|
||||
func NewCachedDNSResolver(resolver DNSResolver) *CachedDNSResolver {
|
||||
return &CachedDNSResolver{
|
||||
resolver: resolver,
|
||||
cache: NewDNSCache(),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *CachedDNSResolver) LookupIP(hostname string) ([]net.IP, error) {
|
||||
if ips, exists := c.cache.Get(hostname); exists {
|
||||
return ips, nil
|
||||
if len(c.data) >= maxDNSPinnedHosts {
|
||||
now := time.Now()
|
||||
for host, entry := range c.data {
|
||||
if now.After(entry.expiresAt) {
|
||||
delete(c.data, host)
|
||||
}
|
||||
}
|
||||
for host := range c.data {
|
||||
if len(c.data) < maxDNSPinnedHosts {
|
||||
break
|
||||
}
|
||||
delete(c.data, host)
|
||||
}
|
||||
}
|
||||
|
||||
ips, err := c.resolver.LookupIP(hostname)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
c.data[hostname] = dnsCacheEntry{
|
||||
ips: ips,
|
||||
expiresAt: time.Now().Add(dnsPinTTL),
|
||||
}
|
||||
|
||||
c.cache.Set(hostname, ips)
|
||||
return ips, nil
|
||||
}
|
||||
|
||||
type CustomDialer struct {
|
||||
cache *DNSCache
|
||||
fallback *net.Dialer
|
||||
cache *DNSCache
|
||||
dialer *net.Dialer
|
||||
}
|
||||
|
||||
func NewCustomDialer(cache *DNSCache) *CustomDialer {
|
||||
return &CustomDialer{
|
||||
cache: cache,
|
||||
fallback: &net.Dialer{
|
||||
dialer: &net.Dialer{
|
||||
Timeout: dialTimeout,
|
||||
},
|
||||
}
|
||||
@@ -115,40 +120,43 @@ func NewCustomDialer(cache *DNSCache) *CustomDialer {
|
||||
func (d *CustomDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
host, port, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return d.fallback.DialContext(ctx, network, address)
|
||||
return nil, ErrSSRFBlocked
|
||||
}
|
||||
|
||||
if ips, exists := d.cache.Get(host); exists {
|
||||
for _, ip := range ips {
|
||||
ipAddr := net.JoinHostPort(ip.String(), port)
|
||||
if conn, err := d.fallback.DialContext(ctx, network, ipAddr); err == nil {
|
||||
return conn, nil
|
||||
}
|
||||
ips, exists := d.cache.Get(host)
|
||||
if !exists {
|
||||
return nil, ErrSSRFBlocked
|
||||
}
|
||||
|
||||
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 {
|
||||
client *http.Client
|
||||
resolver DNSResolver
|
||||
dnsCache *DNSCache
|
||||
approvedHosts map[string]bool
|
||||
mu sync.RWMutex
|
||||
client *http.Client
|
||||
resolver DNSResolver
|
||||
dnsCache *DNSCache
|
||||
}
|
||||
|
||||
func NewURLMetadataService() *URLMetadataService {
|
||||
dnsCache := NewDNSCache()
|
||||
cachedResolver := NewCachedDNSResolver(DefaultDNSResolver{})
|
||||
customDialer := NewCustomDialer(dnsCache)
|
||||
|
||||
svc := &URLMetadataService{
|
||||
resolver: cachedResolver,
|
||||
dnsCache: dnsCache,
|
||||
approvedHosts: make(map[string]bool),
|
||||
resolver: DefaultDNSResolver{},
|
||||
dnsCache: NewDNSCache(),
|
||||
}
|
||||
|
||||
customDialer := NewCustomDialer(svc.dnsCache)
|
||||
|
||||
transport := &http.Transport{
|
||||
DialContext: customDialer.DialContext,
|
||||
MaxIdleConns: 100,
|
||||
@@ -166,25 +174,7 @@ func NewURLMetadataService() *URLMetadataService {
|
||||
if len(via) >= maxRedirects {
|
||||
return ErrTooManyRedirects
|
||||
}
|
||||
|
||||
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.validateURLForSSRF(req.URL)
|
||||
},
|
||||
}
|
||||
return svc
|
||||
@@ -204,19 +194,8 @@ func (s *URLMetadataService) FetchTitle(ctx context.Context, rawURL string) (str
|
||||
return "", ErrUnsupportedScheme
|
||||
}
|
||||
|
||||
hostname := parsed.Hostname()
|
||||
s.mu.RLock()
|
||||
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()
|
||||
if err := s.validateURLForSSRF(parsed); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
|
||||
@@ -336,13 +315,20 @@ func (s *URLMetadataService) validateURLForSSRF(u *url.URL) error {
|
||||
return ErrSSRFBlocked
|
||||
}
|
||||
|
||||
ips, err := s.resolver.LookupIP(u.Hostname())
|
||||
if err != nil {
|
||||
hostname := u.Hostname()
|
||||
if _, pinned := s.dnsCache.Get(hostname); pinned {
|
||||
return nil
|
||||
}
|
||||
|
||||
ips, err := s.resolver.LookupIP(hostname)
|
||||
if err != nil || len(ips) == 0 {
|
||||
return ErrSSRFBlocked
|
||||
}
|
||||
if slices.ContainsFunc(ips, isPrivateOrReservedIP) {
|
||||
return ErrSSRFBlocked
|
||||
}
|
||||
|
||||
s.dnsCache.Set(hostname, ips)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -361,93 +347,42 @@ func isLocalhost(hostname string) bool {
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
ipv4 := ip.To4()
|
||||
if ipv4 == nil {
|
||||
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) {
|
||||
for _, prefix := range reservedPrefixes {
|
||||
if prefix.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user