Add pkg/ssrf shared egress guard and use it for WithNoLocalIP
Centralizes SSRF egress protection so detectors and downstream consumers (the thog analyzer guard) share one classification and one enforcement mechanism. pkg/ssrf provides ErrEgressBlocked, IsNonPublicIP, CheckDialAddress, and GuardDialer. The blocked set is broader than the old isLocalIP: it adds CGNAT (incl. Alibaba metadata), class E, test-nets and other special-use ranges, all multicast, and IPv6 embeddings of internal v4 targets (NAT64, 6to4, IPv4-compatible) that To4() does not normalize. WithNoLocalIP now installs a dialer whose Control hook checks the resolved address immediately before connect. The previous implementation vetted a LookupIP result and then dialed the hostname again, which a DNS-rebinding resolver could race; Control runs on the address actually being dialed, and re-runs on every redirect hop. ErrNoLocalIP remains as a deprecated alias of ssrf.ErrEgressBlocked so errors.Is call sites keep working.
This commit is contained in:
+24
-33
@@ -3,12 +3,10 @@ package detectors
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"slices"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -16,6 +14,7 @@ import (
|
||||
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/common"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/feature"
|
||||
"github.com/trufflesecurity/trufflehog/v3/pkg/ssrf"
|
||||
)
|
||||
|
||||
var DetectorHttpClientWithNoLocalAddresses *http.Client
|
||||
@@ -98,16 +97,26 @@ func NewDetectorTransport(T http.RoundTripper) http.RoundTripper {
|
||||
return &detectorTransport{T: T}
|
||||
}
|
||||
|
||||
func isLocalIP(ip net.IP) bool {
|
||||
if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsPrivate() || ip.IsUnspecified() {
|
||||
return true
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
var ErrNoLocalIP = errors.New("dialing local IP addresses is not allowed")
|
||||
// ErrNoLocalIP is returned when a client configured with WithNoLocalIP
|
||||
// refuses to dial a non-public address.
|
||||
//
|
||||
// Deprecated: it is an alias of ssrf.ErrEgressBlocked; new code should use
|
||||
// that sentinel directly.
|
||||
var ErrNoLocalIP = ssrf.ErrEgressBlocked
|
||||
|
||||
// WithNoLocalIP configures the client to refuse connections to non-public
|
||||
// addresses: loopback, link-local (incl. cloud metadata), private, CGNAT,
|
||||
// multicast, and the other special-use ranges the ssrf package classifies.
|
||||
// The check runs in the guarded dialer on each resolved address, so it is
|
||||
// DNS-rebinding safe (the previous implementation vetted a LookupIP result
|
||||
// and then dialed the hostname again, which a rebinding resolver could race)
|
||||
// and it re-runs on every redirect hop's fresh dial.
|
||||
//
|
||||
// Two behavior notes against the previous implementation: any DialContext
|
||||
// already set on the transport is REPLACED, not chained, so custom dialers
|
||||
// (SOCKS, custom resolvers) are discarded; and blocking is per resolved
|
||||
// address, so a hostname resolving to both a public and a non-public record
|
||||
// connects via the public one where it previously refused the whole host.
|
||||
func WithNoLocalIP() ClientOption {
|
||||
return func(c *http.Client) {
|
||||
if c.Transport == nil {
|
||||
@@ -128,28 +137,10 @@ func WithNoLocalIP() ClientOption {
|
||||
}
|
||||
}
|
||||
|
||||
// If the original DialContext is nil, set it to the default dialer
|
||||
if transport.DialContext == nil {
|
||||
transport.DialContext = defaultDialer.DialContext
|
||||
}
|
||||
originalDialContext := transport.DialContext
|
||||
transport.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
host, port, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ips, err := net.LookupIP(host)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if slices.ContainsFunc(ips, isLocalIP) {
|
||||
return nil, ErrNoLocalIP
|
||||
}
|
||||
|
||||
return originalDialContext(ctx, network, net.JoinHostPort(host, port))
|
||||
}
|
||||
// The guarded dialer replaces any existing DialContext; within this
|
||||
// package that is only ever nil or defaultDialer, whose settings the
|
||||
// guard preserves.
|
||||
transport.DialContext = ssrf.GuardDialer(defaultDialer).DialContext
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -3,7 +3,6 @@ package detectors
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
@@ -308,27 +307,3 @@ func TestDoWithDedup_DeadlinePreserved(t *testing.T) {
|
||||
assert.Error(t, err, "request to hanging server should fail")
|
||||
assert.Less(t, elapsed, time.Second, "timeout should be enforced by client deadline, not run indefinitely")
|
||||
}
|
||||
|
||||
func TestIsLocalIP(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
ip net.IP
|
||||
expected bool
|
||||
}{
|
||||
{"Loopback IPv4", net.ParseIP("127.0.0.1"), true},
|
||||
{"Loopback IPv6", net.ParseIP("::1"), true},
|
||||
{"Private IPv4", net.ParseIP("192.168.1.1"), true},
|
||||
{"Private IPv6", net.ParseIP("fd00::1"), true},
|
||||
{"Unspecified IPv4", net.ParseIP("0.0.0.0"), true},
|
||||
{"Unspecified IPv6", net.ParseIP("::"), true},
|
||||
{"Public IPv4", net.ParseIP("8.8.8.8"), false},
|
||||
{"Public IPv6", net.ParseIP("2001:4860:4860::8888"), false},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
result := isLocalIP(tc.ip)
|
||||
assert.Equal(t, tc.expected, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
// Package ssrf guards outbound connections whose destinations are taken from
|
||||
// scanned or otherwise untrusted content (a secret's "domain"/"endpoint",
|
||||
// a connection string, a URL in a config file). Without a guard, an attacker
|
||||
// who can plant such content can steer an HTTP client at internal-only
|
||||
// addresses, and redirects turn an attacker-owned endpoint into a pivot into
|
||||
// internal space (cluster services, the cloud metadata server, etc.).
|
||||
//
|
||||
// The check runs inside the dialer (via ControlContext), after DNS
|
||||
// resolution and immediately before the socket connects. That placement
|
||||
// matters for two reasons:
|
||||
// - It is DNS-rebinding safe: a hostname that resolves to a public IP on a
|
||||
// pre-flight check but to 169.254.169.254 at connect time is still caught,
|
||||
// because the check runs on the address actually being dialed.
|
||||
// - It covers HTTP redirects: any redirect hop that opens a new connection
|
||||
// re-dials through the guarded dialer, and every fresh dial re-checks the
|
||||
// resolved IP.
|
||||
//
|
||||
// Caveat: the guard inspects the address the transport dials. If an HTTP
|
||||
// forward proxy is configured (HTTP(S)_PROXY), the transport dials the proxy
|
||||
// and the proxy makes the final connection, so egress policy must also be
|
||||
// enforced at the proxy.
|
||||
package ssrf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// ErrEgressBlocked is wrapped by every error this package returns when it
|
||||
// refuses a connection. Callers can errors.Is against it to treat a blocked
|
||||
// egress as an expected, benign outcome rather than a failure.
|
||||
var ErrEgressBlocked = errors.New("egress blocked: non-public address")
|
||||
|
||||
// GuardDialer returns a copy of base (or a zero dialer when base is nil)
|
||||
// that refuses non-public targets via CheckDialAddress. The check is
|
||||
// installed as ControlContext because the net package ignores Control
|
||||
// whenever ControlContext is set: installing the guard as Control would let
|
||||
// a base dialer carrying a ControlContext bypass it entirely. Any
|
||||
// ControlContext or Control already set on base runs after the check, for
|
||||
// allowed targets only.
|
||||
func GuardDialer(base *net.Dialer) *net.Dialer {
|
||||
d := net.Dialer{}
|
||||
if base != nil {
|
||||
d = *base
|
||||
}
|
||||
prevCtx, prevCtl := d.ControlContext, d.Control
|
||||
d.Control = nil
|
||||
d.ControlContext = func(ctx context.Context, network, address string, c syscall.RawConn) error {
|
||||
if err := CheckDialAddress(address); err != nil {
|
||||
return err
|
||||
}
|
||||
if prevCtx != nil {
|
||||
return prevCtx(ctx, network, address, c)
|
||||
}
|
||||
if prevCtl != nil {
|
||||
return prevCtl(network, address, c)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return &d
|
||||
}
|
||||
|
||||
// CheckDialAddress rejects a dial target whose IP is not a public address.
|
||||
// Every rejection wraps ErrEgressBlocked so callers can recognize it. The
|
||||
// address must be a resolved "ip:port" pair as seen by a dialer Control hook;
|
||||
// anything else is refused (fail closed).
|
||||
func CheckDialAddress(address string) error {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
// Fail closed: if we cannot understand the target, do not connect.
|
||||
// %v (not %w) on the parse error: a second %w would make this a
|
||||
// multi-branch error tree, and whole-tree checks like thog's
|
||||
// IsEgressBlockedOnly would classify the rejection as a real failure.
|
||||
return fmt.Errorf("ssrf guard: cannot parse dial address %q: %v: %w", address, err, ErrEgressBlocked)
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil {
|
||||
// At dial time the address is already resolved to an IP literal. If it
|
||||
// is not, refuse rather than risk connecting to something unvetted.
|
||||
return fmt.Errorf("ssrf guard: dial target %q is not an IP literal: %w", host, ErrEgressBlocked)
|
||||
}
|
||||
if IsNonPublicIP(ip) {
|
||||
return fmt.Errorf("ssrf guard: refusing to connect to non-public address %s: %w", ip, ErrEgressBlocked)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// extraBlockedCIDRs covers ranges not classified by the net.IP helper methods
|
||||
// but that must never be reachable from guarded egress.
|
||||
var extraBlockedCIDRs = mustParseCIDRs(
|
||||
// IPv4 special-use / non-public ranges.
|
||||
"0.0.0.0/8", // RFC1122 "this network"; parts route to localhost on some stacks
|
||||
"100.64.0.0/10", // RFC6598 CGNAT (incl. Alibaba metadata 100.100.100.200)
|
||||
"192.0.0.0/24", // RFC6890 IETF protocol assignments
|
||||
"192.0.2.0/24", // TEST-NET-1
|
||||
"198.18.0.0/15", // RFC2544 benchmarking
|
||||
"198.51.100.0/24", // TEST-NET-2
|
||||
"203.0.113.0/24", // TEST-NET-3
|
||||
"240.0.0.0/4", // RFC1112 class E reserved (also covers 255.255.255.255 broadcast)
|
||||
// IPv6 embeddings of an IPv4 address that To4() does NOT normalize; without
|
||||
// these, an internal v4 target can be smuggled in as an IPv6 literal.
|
||||
"::/96", // RFC4291 IPv4-compatible IPv6 (:: and ::1 are caught earlier)
|
||||
"64:ff9b::/96", // RFC6052 NAT64 well-known prefix
|
||||
"64:ff9b:1::/48", // RFC8215 NAT64 local-use prefix
|
||||
"2002::/16", // RFC3056 6to4
|
||||
// IPv6 special-use parity with the v4 test/doc ranges above.
|
||||
"2001:db8::/32", // RFC3849 documentation
|
||||
"100::/64", // RFC6666 discard-only
|
||||
// Note: IPv4-mapped IPv6 (e.g. ::ffff:169.254.169.254) is handled by the
|
||||
// To4() normalization in IsNonPublicIP, not by a CIDR here — a
|
||||
// ::ffff:0:0/96 entry would match every IPv4 address.
|
||||
)
|
||||
|
||||
// IsNonPublicIP reports whether an IP must not be dialed by a guarded client:
|
||||
// loopback, link-local (incl. cloud metadata), private (RFC1918/RFC4193),
|
||||
// CGNAT, multicast, unspecified, the special-use ranges above, and IPv6
|
||||
// embeddings of any of those.
|
||||
func IsNonPublicIP(ip net.IP) bool {
|
||||
if ip == nil {
|
||||
return true // fail closed
|
||||
}
|
||||
// Normalize IPv4-in-IPv6 so the v4 classification methods apply.
|
||||
if v4 := ip.To4(); v4 != nil {
|
||||
ip = v4
|
||||
}
|
||||
|
||||
if ip.IsLoopback() || // 127.0.0.0/8, ::1
|
||||
ip.IsUnspecified() || // 0.0.0.0, ::
|
||||
ip.IsPrivate() || // RFC1918 (10/8, 172.16/12, 192.168/16) + RFC4193 (fc00::/7)
|
||||
ip.IsLinkLocalUnicast() || // 169.254.0.0/16 (metadata 169.254.169.254, ECS creds 169.254.170.2), fe80::/10
|
||||
ip.IsMulticast() { // 224.0.0.0/4, ff00::/8 (subsumes link-local and interface-local multicast)
|
||||
return true
|
||||
}
|
||||
|
||||
for _, n := range extraBlockedCIDRs {
|
||||
if n.Contains(ip) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func mustParseCIDRs(cidrs ...string) []*net.IPNet {
|
||||
nets := make([]*net.IPNet, 0, len(cidrs))
|
||||
for _, c := range cidrs {
|
||||
_, n, err := net.ParseCIDR(c)
|
||||
if err != nil {
|
||||
panic(fmt.Sprintf("ssrf guard: invalid CIDR %q: %v", c, err))
|
||||
}
|
||||
nets = append(nets, n)
|
||||
}
|
||||
return nets
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
package ssrf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"syscall"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIsNonPublicIP(t *testing.T) {
|
||||
cases := []struct {
|
||||
ip string
|
||||
blocked bool
|
||||
}{
|
||||
// Public: allowed.
|
||||
{"8.8.8.8", false},
|
||||
{"1.1.1.1", false},
|
||||
{"93.184.216.34", false}, // example.com
|
||||
{"2606:2800:220:1:248:1893:25c8:1946", false},
|
||||
|
||||
// Loopback.
|
||||
{"127.0.0.1", true},
|
||||
{"127.1.2.3", true},
|
||||
{"::1", true},
|
||||
|
||||
// Link-local incl. cloud metadata and ECS credentials endpoint.
|
||||
{"169.254.169.254", true}, // GCP/AWS/Azure metadata
|
||||
{"169.254.170.2", true}, // AWS ECS task role credentials
|
||||
{"fe80::1", true},
|
||||
|
||||
// Private RFC1918 / RFC4193.
|
||||
{"10.0.0.5", true},
|
||||
{"172.16.31.9", true},
|
||||
{"192.168.1.1", true},
|
||||
{"fc00::1", true},
|
||||
{"fd12:3456::1", true},
|
||||
|
||||
// CGNAT (incl. Alibaba metadata 100.100.100.200).
|
||||
{"100.64.0.1", true},
|
||||
{"100.100.100.200", true},
|
||||
|
||||
// Unspecified + "this network" 0.0.0.0/8.
|
||||
{"0.0.0.0", true},
|
||||
{"0.1.2.3", true},
|
||||
{"::", true},
|
||||
|
||||
// Class E reserved + broadcast.
|
||||
{"240.0.0.1", true},
|
||||
{"255.255.255.255", true},
|
||||
|
||||
// Test-net / protocol-assignment ranges.
|
||||
{"192.0.2.5", true},
|
||||
{"198.18.0.1", true},
|
||||
{"203.0.113.9", true},
|
||||
|
||||
// IPv4-mapped IPv6 must not bypass the v4 checks.
|
||||
{"::ffff:169.254.169.254", true},
|
||||
{"::ffff:10.0.0.1", true},
|
||||
{"::ffff:8.8.8.8", false},
|
||||
|
||||
// IPv6 embeddings of an internal v4 target (To4() does not normalize these).
|
||||
{"64:ff9b::a00:1", true}, // NAT64 of 10.0.0.1
|
||||
{"2002:0a00:0001::1", true}, // 6to4 of 10.0.0.1
|
||||
{"::0a00:0001", true}, // IPv4-compatible ::10.0.0.1
|
||||
{"2001:db8::1", true}, // documentation
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
ip := net.ParseIP(c.ip)
|
||||
require.NotNilf(t, ip, "bad test IP %q", c.ip)
|
||||
assert.Equalf(t, c.blocked, IsNonPublicIP(ip), "IsNonPublicIP(%s)", c.ip)
|
||||
}
|
||||
|
||||
assert.True(t, IsNonPublicIP(nil), "nil IP must fail closed")
|
||||
}
|
||||
|
||||
func TestCheckDialAddress(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
address string
|
||||
blocked bool
|
||||
}{
|
||||
{"public IP allowed", "8.8.8.8:443", false},
|
||||
{"loopback blocked", "127.0.0.1:8080", true},
|
||||
{"metadata blocked", "169.254.169.254:80", true},
|
||||
{"IPv6 loopback blocked", "[::1]:443", true},
|
||||
{"unparseable address blocked", "no-port-here", true},
|
||||
{"unresolved hostname blocked", "example.com:443", true},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
err := CheckDialAddress(c.address)
|
||||
if !c.blocked {
|
||||
assert.NoError(t, err)
|
||||
return
|
||||
}
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, ErrEgressBlocked)
|
||||
assert.Contains(t, err.Error(), "ssrf guard")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardDialer_BlocksLoopback drives a real dial through the guarded dialer
|
||||
// at a loopback httptest server and verifies the refusal carries the sentinel.
|
||||
func TestGuardDialer_BlocksLoopback(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
dialer := GuardDialer(nil)
|
||||
_, err := dialer.DialContext(context.Background(), "tcp", strings.TrimPrefix(srv.URL, "http://"))
|
||||
require.Error(t, err)
|
||||
assert.ErrorIs(t, err, ErrEgressBlocked)
|
||||
}
|
||||
|
||||
// TestGuardDialer_ChainsControl verifies an existing Control hook on the base
|
||||
// dialer still runs (after the guard check) for allowed targets.
|
||||
func TestGuardDialer_ChainsControl(t *testing.T) {
|
||||
var chained bool
|
||||
base := &net.Dialer{
|
||||
Control: func(_, _ string, _ syscall.RawConn) error {
|
||||
chained = true
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Blocked target: the guard refuses before the chained Control runs.
|
||||
dialer := GuardDialer(base)
|
||||
_, err := dialer.DialContext(ctx, "tcp", "127.0.0.1:1")
|
||||
require.ErrorIs(t, err, ErrEgressBlocked)
|
||||
assert.False(t, chained, "chained Control must not run for blocked targets")
|
||||
|
||||
// Allowed target: the chained Control runs. The dial itself may fail
|
||||
// (network-dependent); only the Control invocation matters here.
|
||||
_, _ = dialer.DialContext(ctx, "tcp", "1.1.1.1:53")
|
||||
assert.True(t, chained, "chained Control must run for allowed targets")
|
||||
}
|
||||
|
||||
// TestGuardDialer_GuardsBaseWithControlContext locks in the fix for a bypass:
|
||||
// the net package ignores Control whenever ControlContext is set, so a guard
|
||||
// installed as Control would never run for a base dialer carrying a
|
||||
// ControlContext. The guard must win regardless of which hook the base uses.
|
||||
func TestGuardDialer_GuardsBaseWithControlContext(t *testing.T) {
|
||||
var chained bool
|
||||
base := &net.Dialer{
|
||||
ControlContext: func(_ context.Context, _, _ string, _ syscall.RawConn) error {
|
||||
chained = true
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
dialer := GuardDialer(base)
|
||||
_, err := dialer.DialContext(ctx, "tcp", "127.0.0.1:1")
|
||||
require.ErrorIs(t, err, ErrEgressBlocked,
|
||||
"a base ControlContext must not bypass the guard")
|
||||
assert.False(t, chained, "chained ControlContext must not run for blocked targets")
|
||||
|
||||
// Allowed target: the base's ControlContext still runs after the check.
|
||||
_, _ = dialer.DialContext(ctx, "tcp", "1.1.1.1:53")
|
||||
assert.True(t, chained, "chained ControlContext must run for allowed targets")
|
||||
}
|
||||
Reference in New Issue
Block a user