diff --git a/pkg/registry/transport.go b/pkg/registry/transport.go index fa0103ac0..0eba9b469 100644 --- a/pkg/registry/transport.go +++ b/pkg/registry/transport.go @@ -23,9 +23,11 @@ import ( "io" "log/slog" "mime" + "net" "net/http" "strings" "sync/atomic" + "time" "oras.land/oras-go/v2/registry/remote/retry" ) @@ -53,12 +55,7 @@ type LoggingTransport struct { // NewTransport creates and returns a new instance of LoggingTransport func NewTransport(debug bool) *retry.Transport { - // clone http.DefaultTransport so mutations (e.g. TLS config) are not - // reflected globally - transport := http.DefaultTransport - if t, ok := transport.(*http.Transport); ok { - transport = t.Clone() - } + transport := defaultTransport() if debug { transport = &LoggingTransport{RoundTripper: transport} } @@ -66,6 +63,35 @@ func NewTransport(debug bool) *retry.Transport { return retry.NewTransport(transport) } +// defaultTransport returns the transport a registry client starts from. It is +// always a value Helm owns, because TLS settings are applied to it afterwards +// (see ensureTLSConfig) and must not be applied to the shared +// http.DefaultTransport. +func defaultTransport() http.RoundTripper { + if t, ok := http.DefaultTransport.(*http.Transport); ok { + return t.Clone() + } + + // http.DefaultTransport is a variable of interface type, so an imported package + // can replace it with a RoundTripper that cannot be cloned and cannot carry a + // TLS configuration. Use the standard library defaults rather than handing out + // that replacement, which would make every TLS option fail to apply. + slog.Warn("http.DefaultTransport is not an *http.Transport, using the default transport settings instead", + "transport", fmt.Sprintf("%T", http.DefaultTransport)) + return &http.Transport{ + Proxy: http.ProxyFromEnvironment, + DialContext: (&net.Dialer{ + Timeout: 30 * time.Second, + KeepAlive: 30 * time.Second, + }).DialContext, + ForceAttemptHTTP2: true, + MaxIdleConns: 100, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 10 * time.Second, + ExpectContinueTimeout: 1 * time.Second, + } +} + // RoundTrip calls base round trip while keeping track of the current request. func (t *LoggingTransport) RoundTrip(req *http.Request) (resp *http.Response, err error) { id := requestCount.Add(1) - 1 diff --git a/pkg/registry/transport_test.go b/pkg/registry/transport_test.go index e72c97c0d..ab7b73b48 100644 --- a/pkg/registry/transport_test.go +++ b/pkg/registry/transport_test.go @@ -18,6 +18,7 @@ package registry import ( "bytes" + "crypto/tls" "errors" "io" "net/http" @@ -25,6 +26,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "oras.land/oras-go/v2/registry/remote/auth" + "oras.land/oras-go/v2/registry/remote/retry" ) var errMockRead = errors.New("mock read error") @@ -35,6 +38,24 @@ func (e *errorReader) Read(_ []byte) (n int, err error) { return 0, errMockRead } +// nonClonableTransport stands in for a RoundTripper an imported package replaced +// http.DefaultTransport with: DefaultTransport is a variable of interface type, so +// it does not always hold an *http.Transport. +type nonClonableTransport struct{} + +func (*nonClonableTransport) RoundTrip(*http.Request) (*http.Response, error) { + return nil, errors.New("never called") +} + +// baseTransport unwraps the transport NewTransport built the retry transport from. +func baseTransport(t *testing.T, transport *retry.Transport) http.RoundTripper { + t.Helper() + if logged, ok := transport.Base.(*LoggingTransport); ok { + return logged.RoundTripper + } + return transport.Base +} + func Test_isPrintableContentType(t *testing.T) { tests := []struct { name string @@ -384,3 +405,47 @@ func Test_containsCredentials(t *testing.T) { }) } } + +func Test_NewTransport(t *testing.T) { + defer func(orig http.RoundTripper) { http.DefaultTransport = orig }(http.DefaultTransport) + + tests := []struct { + name string + debug bool + installed http.RoundTripper + }{ + {name: "clonable default transport", installed: &http.Transport{}}, + {name: "replaced default transport", installed: &nonClonableTransport{}}, + {name: "clonable default transport, debug", debug: true, installed: &http.Transport{}}, + {name: "replaced default transport, debug", debug: true, installed: &nonClonableTransport{}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + http.DefaultTransport = tt.installed + + transport := NewTransport(tt.debug) + base := baseTransport(t, transport) + + if tt.debug { + _, ok := transport.Base.(*LoggingTransport) + require.True(t, ok, "expected the transport to be wrapped for debug output, got %T", transport.Base) + } + + assert.NotSame(t, tt.installed, base, "the process-wide default transport must not be handed out") + + httpTransport, ok := base.(*http.Transport) + require.True(t, ok, "TLS settings can only be applied to an *http.Transport, got %T", base) + + conf := &tls.Config{InsecureSkipVerify: true} + _, err := ensureTLSConfig(&auth.Client{Client: &http.Client{Transport: transport}}, conf) + require.NoError(t, err, "TLS settings should be applicable to a Helm registry transport") + assert.Same(t, conf, httpTransport.TLSClientConfig) + + if shared, ok := tt.installed.(*http.Transport); ok { + assert.NotSame(t, conf, shared.TLSClientConfig, + "configuring Helm's transport must not reach the process-wide default transport") + } + }) + } +}