TUN-10701: Use curves for prechecks

Use curves when running pre-checks. Although it is not something critical, pre-checks should closely match the current cloudflared behavior when trying to establish connections to the edge. Adding curves here matches the current behavor.
This commit is contained in:
Miguel da Costa Martins Marcelino
2026-07-20 10:00:35 +00:00
parent 8679787525
commit 2206516c3b
6 changed files with 48 additions and 31 deletions
+1 -1
View File
@@ -142,7 +142,7 @@ func installLaunchd(c *cli.Context) error {
etPath, err := os.Executable() etPath, err := os.Executable()
if err != nil { if err != nil {
log.Err(err).Msg("Error determining executable path") log.Err(err).Msg("Error determining executable path")
return fmt.Errorf("Error determining executable path: %w", err) return fmt.Errorf("error determining executable path: %w", err)
} }
installPath, err := installPath() installPath, err := installPath()
if err != nil { if err != nil {
+4 -3
View File
@@ -33,6 +33,7 @@ import (
"github.com/cloudflare/cloudflared/diagnostic" "github.com/cloudflare/cloudflared/diagnostic"
"github.com/cloudflare/cloudflared/edgediscovery" "github.com/cloudflare/cloudflared/edgediscovery"
"github.com/cloudflare/cloudflared/edgediscovery/allregions" "github.com/cloudflare/cloudflared/edgediscovery/allregions"
"github.com/cloudflare/cloudflared/features"
"github.com/cloudflare/cloudflared/ingress" "github.com/cloudflare/cloudflared/ingress"
"github.com/cloudflare/cloudflared/logger" "github.com/cloudflare/cloudflared/logger"
"github.com/cloudflare/cloudflared/management" "github.com/cloudflare/cloudflared/management"
@@ -421,7 +422,7 @@ func StartServer(
// goroutine, as we want to keep initializing cloudflared while prechecks // goroutine, as we want to keep initializing cloudflared while prechecks
// are running. Prechecks are controlled via DNS flag for remote kill-switch capability. // are running. Prechecks are controlled via DNS flag for remote kill-switch capability.
if !tunnelConfig.ClientConfig.ConnectionFeaturesSnapshot().SkipPrechecks && !c.Bool(cfdflags.NoPrechecks) { if !tunnelConfig.ClientConfig.ConnectionFeaturesSnapshot().SkipPrechecks && !c.Bool(cfdflags.NoPrechecks) {
go runPrechecks(c, log, tunnelConfig.Region) go runPrechecks(c, log, tunnelConfig.Region, tunnelConfig.ClientConfig.ConnectionFeaturesSnapshot().PostQuantum)
} }
// Disable ICMP packet routing for quick tunnels // Disable ICMP packet routing for quick tunnels
@@ -525,7 +526,7 @@ func StartServer(
// runPrechecks executes connectivity pre-checks and logs the results. // runPrechecks executes connectivity pre-checks and logs the results.
// Pre-checks are diagnostic only and do not gate tunnel startup. // Pre-checks are diagnostic only and do not gate tunnel startup.
func runPrechecks(c *cli.Context, log *zerolog.Logger, region string) { func runPrechecks(c *cli.Context, log *zerolog.Logger, region string, pqMode features.PostQuantumMode) {
ipVersion := allregions.Auto ipVersion := allregions.Auto
if ipVersionStr := c.String(cfdflags.EdgeIpVersion); ipVersionStr != "" { if ipVersionStr := c.String(cfdflags.EdgeIpVersion); ipVersionStr != "" {
parsedVersion, err := parseConfigIPVersion(ipVersionStr) parsedVersion, err := parseConfigIPVersion(ipVersionStr)
@@ -550,7 +551,7 @@ func runPrechecks(c *cli.Context, log *zerolog.Logger, region string) {
ManagementDialer: &prechecks.NetManagementDialer{Dialer: net.Dialer{}}, ManagementDialer: &prechecks.NetManagementDialer{Dialer: net.Dialer{}},
} }
report := prechecks.Run(c.Context, c.String(cfdflags.CACert), cfg, log, dialers) report := prechecks.Run(c.Context, c.String(cfdflags.CACert), cfg, pqMode, log, dialers)
// Output the human-readable table // Output the human-readable table
cliutil.LogTable(log, report.String(), "CONNECTIVITY PRE-CHECKS") cliutil.LogTable(log, report.String(), "CONNECTIVITY PRE-CHECKS")
+2 -1
View File
@@ -18,6 +18,7 @@ import (
network "github.com/cloudflare/cloudflared/diagnostic/network" network "github.com/cloudflare/cloudflared/diagnostic/network"
"github.com/cloudflare/cloudflared/edgediscovery/allregions" "github.com/cloudflare/cloudflared/edgediscovery/allregions"
"github.com/cloudflare/cloudflared/features"
"github.com/cloudflare/cloudflared/prechecks" "github.com/cloudflare/cloudflared/prechecks"
) )
@@ -470,7 +471,7 @@ func collectPrechecks(region string) collectFunc {
} }
emptyCert := "" emptyCert := ""
report := prechecks.Run(ctx, emptyCert, cfg, &log, dialers) report := prechecks.Run(ctx, emptyCert, cfg, features.PostQuantumPrefer, &log, dialers)
// Write the report to a JSON file // Write the report to a JSON file
// nolint: gosec // nolint: gosec
+9 -4
View File
@@ -12,6 +12,7 @@ import (
"github.com/cloudflare/cloudflared/connection" "github.com/cloudflare/cloudflared/connection"
"github.com/cloudflare/cloudflared/edgediscovery/allregions" "github.com/cloudflare/cloudflared/edgediscovery/allregions"
"github.com/cloudflare/cloudflared/features"
) )
const ( const (
@@ -59,7 +60,10 @@ func (tr TransportResults) Collect() []CheckResult {
// //
// Each failed probe is retried up to maxRetries times with exponential backoff. // Each failed probe is retried up to maxRetries times with exponential backoff.
// The suite is bounded by cfg.Timeout (defaultTimeout if zero). // The suite is bounded by cfg.Timeout (defaultTimeout if zero).
func Run(ctx context.Context, caCert string, cfg Config, log *zerolog.Logger, runDialers RunDialers) Report { //
// pqMode controls the TLS curve preferences advertised during probe handshakes,
// matching the key-exchange algorithms used by the real tunnel connections.
func Run(ctx context.Context, caCert string, cfg Config, pqMode features.PostQuantumMode, log *zerolog.Logger, runDialers RunDialers) Report {
runID := uuid.New() runID := uuid.New()
if cfg.Timeout <= 0 { if cfg.Timeout <= 0 {
@@ -68,9 +72,10 @@ func Run(ctx context.Context, caCert string, cfg Config, log *zerolog.Logger, ru
ctx, cancel := context.WithTimeout(ctx, cfg.Timeout) ctx, cancel := context.WithTimeout(ctx, cfg.Timeout)
defer cancel() defer cancel()
// Build TLS configs once per protocol. // Build TLS configs once per protocol, applying the same curve preferences
quicTLSConfig, quicTLSErr := probeTLSConfig(caCert, connection.QUIC) // (including post-quantum curves) used by production tunnel connections.
http2TLSConfig, http2TLSErr := probeTLSConfig(caCert, connection.HTTP2) quicTLSConfig, quicTLSErr := probeTLSConfig(caCert, connection.QUIC, pqMode)
http2TLSConfig, http2TLSErr := probeTLSConfig(caCert, connection.HTTP2, pqMode)
// 1) Resolve edge addresses. Each ResolvedTarget bundles its addr group // 1) Resolve edge addresses. Each ResolvedTarget bundles its addr group
// with the DNS CheckResult that labels it, keeping the two in sync. // with the DNS CheckResult that labels it, keeping the two in sync.
+19 -18
View File
@@ -16,6 +16,7 @@ import (
"github.com/cloudflare/cloudflared/connection" "github.com/cloudflare/cloudflared/connection"
"github.com/cloudflare/cloudflared/edgediscovery/allregions" "github.com/cloudflare/cloudflared/edgediscovery/allregions"
"github.com/cloudflare/cloudflared/features"
"github.com/cloudflare/cloudflared/mocks" "github.com/cloudflare/cloudflared/mocks"
) )
@@ -119,7 +120,7 @@ func TestRun_AllPass(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS + 2 QUIC + 2 HTTP2 + 1 API = 7 results. // 2 DNS + 2 QUIC + 2 HTTP2 + 1 API = 7 results.
requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Pass) requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Pass)
@@ -150,7 +151,7 @@ func TestRun_QUICBlocked(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS Pass + 2 QUIC Fail + 2 HTTP2 Pass + 1 API Pass. // 2 DNS Pass + 2 QUIC Fail + 2 HTTP2 Pass + 1 API Pass.
requireStatuses(t, report, Pass, Pass, Fail, Fail, Pass, Pass, Pass) requireStatuses(t, report, Pass, Pass, Fail, Fail, Pass, Pass, Pass)
@@ -180,7 +181,7 @@ func TestRun_HTTP2Blocked(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS Pass + 2 QUIC Pass + 2 HTTP2 Fail + 1 API Pass. // 2 DNS Pass + 2 QUIC Pass + 2 HTTP2 Fail + 1 API Pass.
requireStatuses(t, report, Pass, Pass, Pass, Pass, Fail, Fail, Pass) requireStatuses(t, report, Pass, Pass, Pass, Pass, Fail, Fail, Pass)
@@ -210,7 +211,7 @@ func TestRun_BothTransportsBlocked(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS Pass + 2 QUIC Fail + 2 HTTP2 Fail + 1 API Pass. // 2 DNS Pass + 2 QUIC Fail + 2 HTTP2 Fail + 1 API Pass.
requireStatuses(t, report, Pass, Pass, Fail, Fail, Fail, Fail, Pass) requireStatuses(t, report, Pass, Pass, Fail, Fail, Fail, Fail, Pass)
@@ -249,7 +250,7 @@ func TestRun_PartialRegionQUICFail(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS Pass + QUIC-region1 Pass + QUIC-region2 Fail + 2 HTTP2 Pass + 1 API Pass. // 2 DNS Pass + QUIC-region1 Pass + QUIC-region2 Fail + 2 HTTP2 Pass + 1 API Pass.
requireStatuses(t, report, Pass, Pass, Pass, Fail, Pass, Pass, Pass) requireStatuses(t, report, Pass, Pass, Pass, Fail, Pass, Pass, Pass)
@@ -282,7 +283,7 @@ func TestRun_DNSFail_SkipsTransports(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// DNS failure emits 2 Fail rows (one per default region). // DNS failure emits 2 Fail rows (one per default region).
// Transport rows: one skip per DNS region for QUIC and HTTP/2 = 2 QUIC skips + 2 HTTP2 skips. // Transport rows: one skip per DNS region for QUIC and HTTP/2 = 2 QUIC skips + 2 HTTP2 skips.
@@ -319,7 +320,7 @@ func TestRun_ManagementAPIFail(t *testing.T) {
Return(nil, errors.New("connection refused")).AnyTimes() Return(nil, errors.New("connection refused")).AnyTimes()
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS Pass + 2 QUIC Pass + 2 HTTP2 Pass + 1 API Fail. // 2 DNS Pass + 2 QUIC Pass + 2 HTTP2 Pass + 1 API Fail.
requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Fail) requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Fail)
@@ -350,7 +351,7 @@ func TestRun_RegionFlagForwardedToDNS(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Region: "us", Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Region: "us", Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// DNS rows carry regional hostnames (indices 0 and 1). // DNS rows carry regional hostnames (indices 0 and 1).
assert.Equal(t, "us-region1.v2.argotunnel.com", report.Results[0].Target, "DNS region1") assert.Equal(t, "us-region1.v2.argotunnel.com", report.Results[0].Target, "DNS region1")
@@ -388,7 +389,7 @@ func TestRun_QUICUsesProbeConnIndex(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
} }
// TestRun_BothFamiliesProbed verifies that when both V4 and V6 addresses are // TestRun_BothFamiliesProbed verifies that when both V4 and V6 addresses are
@@ -412,7 +413,7 @@ func TestRun_BothFamiliesProbed(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: allregions.Auto},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS + 2 QUIC + 2 HTTP2 + 1 API = 7 results, all passing. // 2 DNS + 2 QUIC + 2 HTTP2 + 1 API = 7 results, all passing.
requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Pass) requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Pass)
@@ -454,7 +455,7 @@ func TestRun_IPVersionRestriction(t *testing.T) {
Return(nopConn{}, nil) Return(nopConn{}, nil)
report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: tt.ipVersion}, report := Run(t.Context(), emptyCert, Config{Timeout: 2 * time.Second, IPVersion: tt.ipVersion},
nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) features.PostQuantumPrefer, nopLogger(), RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Pass) requireStatuses(t, report, Pass, Pass, Pass, Pass, Pass, Pass, Pass)
}) })
@@ -489,7 +490,7 @@ func TestRun_EdgeAddrs_SingleAddr(t *testing.T) {
Timeout: 2 * time.Second, Timeout: 2 * time.Second,
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 1 DNS Skip + 1 QUIC + 1 HTTP2 + 1 API = 4 results. // 1 DNS Skip + 1 QUIC + 1 HTTP2 + 1 API = 4 results.
@@ -527,7 +528,7 @@ func TestRun_EdgeAddrs_MultipleAddrs(t *testing.T) {
Timeout: 2 * time.Second, Timeout: 2 * time.Second,
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 2 DNS Pass (one per addr) + 2 QUIC + 2 HTTP2 + 1 API = 7 results. // 2 DNS Pass (one per addr) + 2 QUIC + 2 HTTP2 + 1 API = 7 results.
@@ -567,7 +568,7 @@ func TestRun_EdgeAddrs_UnresolvableAddr(t *testing.T) {
Timeout: 2 * time.Second, Timeout: 2 * time.Second,
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// 1 DNS Fail + 1 QUIC Skip + 1 HTTP2 Skip + 1 API = 4 results. // 1 DNS Fail + 1 QUIC Skip + 1 HTTP2 Skip + 1 API = 4 results.
@@ -609,7 +610,7 @@ func TestRun_ProtocolOverride_HTTP2_BothPass(t *testing.T) {
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
ProtocolOverride: "http2", ProtocolOverride: "http2",
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// Both transports pass, but the override must win — HTTP/2 is reported. // Both transports pass, but the override must win — HTTP/2 is reported.
@@ -644,7 +645,7 @@ func TestRun_ProtocolOverride_QUIC_BothPass(t *testing.T) {
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
ProtocolOverride: "quic", ProtocolOverride: "quic",
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
require.NotNil(t, report.SuggestedProtocol) require.NotNil(t, report.SuggestedProtocol)
@@ -677,7 +678,7 @@ func TestRun_ProtocolOverride_HTTP2_QUICBlocked(t *testing.T) {
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
ProtocolOverride: "http2", ProtocolOverride: "http2",
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
require.NotNil(t, report.SuggestedProtocol) require.NotNil(t, report.SuggestedProtocol)
@@ -710,7 +711,7 @@ func TestRun_ProtocolOverride_HTTP2_BothBlocked(t *testing.T) {
IPVersion: allregions.Auto, IPVersion: allregions.Auto,
ProtocolOverride: "http2", ProtocolOverride: "http2",
} }
report := Run(t.Context(), emptyCert, cfg, nopLogger(), report := Run(t.Context(), emptyCert, cfg, features.PostQuantumPrefer, nopLogger(),
RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt}) RunDialers{DNSResolver: dns, TCPDialer: tcp, QUICDialer: quicD, ManagementDialer: mgmt})
// The overridden transport (HTTP/2) is blocked, so the override cannot be // The overridden transport (HTTP/2) is blocked, so the override cannot be
+13 -4
View File
@@ -15,8 +15,10 @@ import (
"github.com/cloudflare/cloudflared/connection/dialopts" "github.com/cloudflare/cloudflared/connection/dialopts"
"github.com/cloudflare/cloudflared/connection" "github.com/cloudflare/cloudflared/connection"
cfdcrypto "github.com/cloudflare/cloudflared/crypto"
edgedial "github.com/cloudflare/cloudflared/edgediscovery" edgedial "github.com/cloudflare/cloudflared/edgediscovery"
"github.com/cloudflare/cloudflared/edgediscovery/allregions" "github.com/cloudflare/cloudflared/edgediscovery/allregions"
"github.com/cloudflare/cloudflared/features"
cfdquic "github.com/cloudflare/cloudflared/quic" cfdquic "github.com/cloudflare/cloudflared/quic"
"github.com/cloudflare/cloudflared/tlsconfig" "github.com/cloudflare/cloudflared/tlsconfig"
) )
@@ -110,10 +112,13 @@ func (d *NetManagementDialer) DialContext(ctx context.Context, network, addr str
} }
// probeTLSConfig builds a *tls.Config for a pre-check probe using the same // probeTLSConfig builds a *tls.Config for a pre-check probe using the same
// certificate pool as the production tunnel. The SNI and NextProtos are taken from // certificate pool and curve preferences as the production tunnel. The SNI and
// p.ProbeTLSSettings() so that the probe SNI is used instead of the production SNI, // NextProtos are taken from p.ProbeTLSSettings() so that the probe SNI is used
// which avoids noisy logs in origintunneld. // instead of the production SNI, which avoids noisy logs in origintunneld.
func probeTLSConfig(caCert string, p connection.Protocol) (*tls.Config, error) { // Curve preferences are set via cfdcrypto.TLSConfigWithCurvePreferences so that
// prechecks advertise the same key-exchange algorithms (including post-quantum
// curves) as the real QUIC/H2 connections.
func probeTLSConfig(caCert string, p connection.Protocol, pqMode features.PostQuantumMode) (*tls.Config, error) {
settings := p.ProbeTLSSettings() settings := p.ProbeTLSSettings()
if settings == nil { if settings == nil {
return nil, fmt.Errorf("no probe TLS settings for protocol %s", p) return nil, fmt.Errorf("no probe TLS settings for protocol %s", p)
@@ -125,6 +130,10 @@ func probeTLSConfig(caCert string, p connection.Protocol) (*tls.Config, error) {
if len(settings.NextProtos) > 0 { if len(settings.NextProtos) > 0 {
cfg.NextProtos = settings.NextProtos cfg.NextProtos = settings.NextProtos
} }
cfg, err = cfdcrypto.TLSConfigWithCurvePreferences(cfg, pqMode)
if err != nil {
return nil, fmt.Errorf("apply curve preferences: %w", err)
}
return cfg, nil return cfg, nil
} }