diff --git a/proxy/internal/proxy/reverseproxy.go b/proxy/internal/proxy/reverseproxy.go index 7b0acd1b3..7583b2e01 100644 --- a/proxy/internal/proxy/reverseproxy.go +++ b/proxy/internal/proxy/reverseproxy.go @@ -809,6 +809,12 @@ func classifyProxyError(err error) (title, message string, code int, status web. http.StatusBadGateway, web.ErrorStatus{Proxy: false, Destination: false} + case errors.Is(err, roundtrip.ErrDirectUpstreamBlocked): + return "Destination Not Allowed", + "This proxy does not connect to private or internal addresses. Please contact your administrator.", + http.StatusBadGateway, + web.ErrorStatus{Proxy: false, Destination: false} + case errors.Is(err, roundtrip.ErrTooManyInflight): return "Service Overloaded", "The service is currently handling too many requests. Please try again shortly.", diff --git a/proxy/internal/proxy/reverseproxy_test.go b/proxy/internal/proxy/reverseproxy_test.go index 83afee387..c0724ce84 100644 --- a/proxy/internal/proxy/reverseproxy_test.go +++ b/proxy/internal/proxy/reverseproxy_test.go @@ -1053,6 +1053,17 @@ func TestClassifyProxyError(t *testing.T) { wantCode: http.StatusBadGateway, wantStatus: web.ErrorStatus{Proxy: true, Destination: false}, }, + { + name: "direct upstream blocked by dial guard", + err: &net.OpError{ + Op: "dial", + Net: "tcp", + Err: roundtrip.ErrDirectUpstreamBlocked, + }, + wantTitle: "Destination Not Allowed", + wantCode: http.StatusBadGateway, + wantStatus: web.ErrorStatus{Proxy: false, Destination: false}, + }, { name: "unknown error falls to default", err: errors.New("something unexpected"), diff --git a/proxy/internal/roundtrip/dialguard.go b/proxy/internal/roundtrip/dialguard.go new file mode 100644 index 000000000..ac01263b4 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard.go @@ -0,0 +1,90 @@ +package roundtrip + +import ( + "context" + "errors" + "net/netip" + "syscall" +) + +// ErrDirectUpstreamBlocked is returned when a direct-upstream dial targets +// an address that is not globally reachable while +// NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE is set. +var ErrDirectUpstreamBlocked = errors.New("direct upstream address is not allowed") + +// blockedUpstreamPrefixes are the ranges that reach the proxy host, its +// cluster or its cloud provider rather than the public internet. NAT64 +// and 6to4 addresses are matched by the IPv4 address they embed. +var blockedUpstreamPrefixes = []netip.Prefix{ + // IPv4 + netip.MustParsePrefix("0.0.0.0/8"), // "this network", including 0.0.0.0 + netip.MustParsePrefix("10.0.0.0/8"), // RFC1918 + netip.MustParsePrefix("100.64.0.0/10"), // CGNAT + netip.MustParsePrefix("127.0.0.0/8"), // loopback + netip.MustParsePrefix("169.254.0.0/16"), // link-local, cloud metadata services + netip.MustParsePrefix("172.16.0.0/12"), // RFC1918 + netip.MustParsePrefix("192.0.0.0/24"), // IETF protocol assignments + netip.MustParsePrefix("192.0.2.0/24"), // documentation + netip.MustParsePrefix("192.88.99.0/24"), // 6to4 relay anycast (deprecated) + netip.MustParsePrefix("192.168.0.0/16"), // RFC1918 + netip.MustParsePrefix("198.18.0.0/15"), // benchmarking + netip.MustParsePrefix("198.51.100.0/24"), // documentation + netip.MustParsePrefix("203.0.113.0/24"), // documentation + netip.MustParsePrefix("224.0.0.0/4"), // multicast + netip.MustParsePrefix("240.0.0.0/4"), // reserved, including broadcast + + // IPv6 + netip.MustParsePrefix("::/96"), // unspecified, loopback, IPv4-compatible + netip.MustParsePrefix("64:ff9b:1::/48"), // local-use NAT64 + netip.MustParsePrefix("100::/64"), // discard-only + netip.MustParsePrefix("2001::/32"), // Teredo + netip.MustParsePrefix("2001:2::/48"), // benchmarking + netip.MustParsePrefix("2001:db8::/32"), // documentation + netip.MustParsePrefix("3fff::/20"), // documentation + netip.MustParsePrefix("5f00::/16"), // SRv6 SIDs + netip.MustParsePrefix("fc00::/7"), // unique local, including AWS IMDS fd00:ec2::254 + netip.MustParsePrefix("fe80::/10"), // link-local + netip.MustParsePrefix("fec0::/10"), // site-local (deprecated) + netip.MustParsePrefix("ff00::/8"), // multicast +} + +var ( + nat64Prefix = netip.MustParsePrefix("64:ff9b::/96") + sixToFour = netip.MustParsePrefix("2002::/16") +) + +// isBlockedUpstreamAddr reports whether a guarded direct-upstream dial +// must refuse addr. +func isBlockedUpstreamAddr(addr netip.Addr) bool { + addr = addr.Unmap().WithZone("") + if !addr.IsValid() { + return true + } + + if nat64Prefix.Contains(addr) { + b := addr.As16() + return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[12:16]))) + } + if sixToFour.Contains(addr) { + b := addr.As16() + return isBlockedUpstreamAddr(netip.AddrFrom4([4]byte(b[2:6]))) + } + + for _, p := range blockedUpstreamPrefixes { + if p.Contains(addr) { + return true + } + } + return false +} + +// guardUpstreamDial is a net.Dialer ControlContext that refuses blocked +// addresses. It sees the resolved address of each socket just before +// connect, so DNS rebinding cannot swap the target after the check. +func guardUpstreamDial(_ context.Context, _, address string, _ syscall.RawConn) error { + ap, err := netip.ParseAddrPort(address) + if err != nil || isBlockedUpstreamAddr(ap.Addr()) { + return ErrDirectUpstreamBlocked + } + return nil +} diff --git a/proxy/internal/roundtrip/dialguard_test.go b/proxy/internal/roundtrip/dialguard_test.go new file mode 100644 index 000000000..79453d8c3 --- /dev/null +++ b/proxy/internal/roundtrip/dialguard_test.go @@ -0,0 +1,189 @@ +package roundtrip + +import ( + "context" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "net/url" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsBlockedUpstreamAddr(t *testing.T) { + blocked := []string{ + "0.0.0.0", + "0.1.2.3", + "10.1.2.3", + "100.64.0.1", + "100.127.255.254", + "127.0.0.1", + "127.255.255.255", + "169.254.169.254", + "172.16.0.1", + "172.31.255.255", + "192.0.0.170", + "192.168.1.1", + "192.88.99.1", + "198.18.0.1", + "224.0.0.1", + "255.255.255.255", + "::", + "::1", + "::169.254.169.254", + "::ffff:127.0.0.1", + "::ffff:169.254.169.254", + "::ffff:10.0.0.1", + "64:ff9b::a9fe:a9fe", // NAT64 of 169.254.169.254 + "64:ff9b::a00:1", // NAT64 of 10.0.0.1 + "64:ff9b:1::1", + "2001::1", + "2001:0:4136:e378:8000:63bf:3fff:fdd2", + "2001:2::1", + "3fff::1", + "5f00::1", + "2002:a9fe:a9fe::1", // 6to4 of 169.254.169.254 + "2002:7f00:1::", // 6to4 of 127.0.0.1 + "fc00::1", + "fd00:ec2::254", + "fe80::1", + "fe80::1%eth0", + "fec0::1", + "ff02::1", + } + for _, s := range blocked { + t.Run("blocks "+s, func(t *testing.T) { + assert.True(t, isBlockedUpstreamAddr(netip.MustParseAddr(s))) + }) + } + + allowed := []string{ + "1.1.1.1", + "8.8.8.8", + "100.63.255.255", + "100.128.0.0", + "172.15.255.255", + "172.32.0.0", + "169.253.255.255", + "2606:4700:4700::1111", + "2001:4860:4860::8888", + "2001:1::1", + "4000::1", + "::ffff:8.8.8.8", + "64:ff9b::808:808", // NAT64 of 8.8.8.8 + "2002:808:808::1", // 6to4 of 8.8.8.8 + } + for _, s := range allowed { + t.Run("allows "+s, func(t *testing.T) { + assert.False(t, isBlockedUpstreamAddr(netip.MustParseAddr(s))) + }) + } + + assert.True(t, isBlockedUpstreamAddr(netip.Addr{}), "the zero Addr must be refused") +} + +func TestGuardUpstreamDial_RejectsUnparsableAddress(t *testing.T) { + err := guardUpstreamDial(context.Background(), "tcp", "not-an-address", nil) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "an address the guard cannot parse must fail closed") +} + +// TestMultiTransport_BlockPrivateUpstreams exercises the guard end to end +// against a loopback test server: by IP literal and by a hostname that +// resolves to loopback, on both direct branches, and confirms the +// embedded branch is not affected. +func TestMultiTransport_BlockPrivateUpstreams(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, "reached") + })) + defer srv.Close() + + _, port, err := net.SplitHostPort(srv.Listener.Addr().String()) + require.NoError(t, err) + byName := (&url.URL{Scheme: "http", Host: net.JoinHostPort("localhost", port)}).String() + + directCtx := WithDirectUpstream(context.Background()) + insecureCtx := WithSkipTLSVerify(directCtx) + + // roundTrip returns the response body, so callers never hold one open. + roundTrip := func(t *testing.T, mt *MultiTransport, ctx context.Context, target string) (string, error) { + t.Helper() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil) + require.NoError(t, err) + resp, err := mt.RoundTrip(req) + if err != nil { + return "", err + } + defer func() { _ = resp.Body.Close() }() + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + return string(body), nil + } + + t.Run("enabled refuses loopback", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "true") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + cases := []struct { + name string + ctx context.Context + target string + }{ + {"direct by IP", directCtx, srv.URL}, + {"direct by hostname", directCtx, byName}, + {"insecure by IP", insecureCtx, srv.URL}, + {"insecure by hostname", insecureCtx, byName}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := roundTrip(t, mt, tc.ctx, tc.target) + require.Error(t, err) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked) + }) + } + }) + + t.Run("invalid value enables the guard", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "yes please") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + _, err := roundTrip(t, mt, directCtx, srv.URL) + assert.ErrorIs(t, err, ErrDirectUpstreamBlocked, "a value that does not parse must fail closed") + }) + + t.Run("explicit false disables the guard", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "false") + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + body, err := roundTrip(t, mt, directCtx, srv.URL) + require.NoError(t, err) + assert.Equal(t, "reached", body) + }) + + t.Run("enabled leaves embedded branch alone", func(t *testing.T) { + t.Setenv(EnvDirectUpstreamBlockPrivate, "true") + embedded := &stubRoundTripper{body: "embedded"} + mt := NewMultiTransport(embedded, nil) + + body, err := roundTrip(t, mt, context.Background(), srv.URL) + require.NoError(t, err) + assert.Equal(t, "embedded", body) + assert.True(t, embedded.called, "the guard must not change dispatch to the embedded transport") + }) + + t.Run("disabled by default", func(t *testing.T) { + // Register the restore first so an exported value comes back after + // the test, then exercise a genuinely absent variable. + t.Setenv(EnvDirectUpstreamBlockPrivate, "") + require.NoError(t, os.Unsetenv(EnvDirectUpstreamBlockPrivate)) + mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil) + + body, err := roundTrip(t, mt, directCtx, srv.URL) + require.NoError(t, err, "private and self-hosted proxies must keep reaching local upstreams") + assert.Equal(t, "reached", body) + }) +} diff --git a/proxy/internal/roundtrip/multi.go b/proxy/internal/roundtrip/multi.go index d50ad1fc9..a430d45bd 100644 --- a/proxy/internal/roundtrip/multi.go +++ b/proxy/internal/roundtrip/multi.go @@ -41,7 +41,9 @@ var errNoEmbeddedTransport = errors.New("multitransport: embedded roundtripper n // MultiTransport that only ever uses the direct branch. The direct // branches honour the same NB_PROXY_* tuning env vars as the embedded // transport (see loadTransportConfig) plus a dial-timeout wrapper that -// respects types.WithDialTimeout. +// respects types.WithDialTimeout. With NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE +// set, the direct branches refuse addresses that are not globally reachable +// (see guardUpstreamDial). func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTransport { if logger == nil { logger = log.StandardLogger() @@ -51,6 +53,9 @@ func NewMultiTransport(embedded http.RoundTripper, logger *log.Logger) *MultiTra Timeout: 30 * time.Second, KeepAlive: 30 * time.Second, } + if cfg.blockPrivateUpstreams { + dialer.ControlContext = guardUpstreamDial + } direct := &http.Transport{ DialContext: dialWithTimeout(dialer.DialContext), MaxIdleConns: cfg.maxIdleConns, diff --git a/proxy/internal/roundtrip/transport.go b/proxy/internal/roundtrip/transport.go index 9e872e447..6383079c4 100644 --- a/proxy/internal/roundtrip/transport.go +++ b/proxy/internal/roundtrip/transport.go @@ -25,6 +25,12 @@ const ( EnvDisableCompression = "NB_PROXY_DISABLE_COMPRESSION" EnvMaxInflight = "NB_PROXY_MAX_INFLIGHT" EnvUpstreamHTTPVersion = "NB_PROXY_UPSTREAM_HTTP_VERSION" + // EnvDirectUpstreamBlockPrivate refuses direct-upstream dials to + // addresses that are not globally reachable (loopback, private, + // link-local, CGNAT, ...). Off by default: private and self-hosted + // proxies use direct_upstream to reach LAN and localhost services. + // Proxies that serve untrusted accounts must turn it on. + EnvDirectUpstreamBlockPrivate = "NB_PROXY_DIRECT_UPSTREAM_BLOCK_PRIVATE" ) // upstreamHTTPVersion selects the HTTP version the proxy uses towards an @@ -69,6 +75,9 @@ type transportConfig struct { // explicit values are for backends whose advertised h2 support is // unusable and whose failure mode the negotiation cannot see. upstreamHTTPVersion upstreamHTTPVersion + // blockPrivateUpstreams guards the direct branches' dialer with + // guardUpstreamDial. It has no effect on the embedded branch. + blockPrivateUpstreams bool } func defaultTransportConfig() transportConfig { @@ -122,6 +131,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig { if v, ok := envUpstreamHTTPVersion(EnvUpstreamHTTPVersion, logger); ok { cfg.upstreamHTTPVersion = v } + cfg.blockPrivateUpstreams = envGuardBool(EnvDirectUpstreamBlockPrivate, logger) logger.WithFields(log.Fields{ "max_idle_conns": cfg.maxIdleConns, @@ -136,6 +146,7 @@ func loadTransportConfig(logger *log.Logger) transportConfig { "disable_compression": cfg.disableCompression, "max_inflight": cfg.maxInflight, "upstream_http_version": cfg.upstreamHTTPVersion, + "block_private_upstreams": cfg.blockPrivateUpstreams, }).Debug("backend transport configuration") return cfg @@ -246,6 +257,22 @@ func envDuration(key string, logger *log.Logger) (time.Duration, bool) { return v, true } +// envGuardBool reads a bool that turns a security guard on. Unset means +// off, but a value that does not parse turns the guard on: a typo must not +// leave a proxy that was meant to be guarded without the guard. +func envGuardBool(key string, logger *log.Logger) bool { + s := os.Getenv(key) + if s == "" { + return false + } + v, err := strconv.ParseBool(s) + if err != nil { + logger.Warnf("failed to parse %s=%q as bool, enabling it: %v", key, s, err) + return true + } + return v +} + func envBool(key string, logger *log.Logger) (bool, bool) { s := os.Getenv(key) if s == "" {