Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 46 additions & 19 deletions oci/common.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
package oci

import (
"crypto/tls"
"fmt"
"log/slog"
"net/http"
Expand All @@ -16,7 +17,6 @@ import (
"oras.land/oras-go/v2/registry/remote"
"oras.land/oras-go/v2/registry/remote/auth"
"oras.land/oras-go/v2/registry/remote/credentials"
"oras.land/oras-go/v2/registry/remote/retry"

"github.com/defenseunicorns/pkg/helpers/v2"
)
Expand All @@ -28,12 +28,13 @@ const (

// OrasRemote is a wrapper around the Oras remote repository that includes a progress bar for interactive feedback.
type OrasRemote struct {
repo *remote.Repository
cache *oci.Store
root *Manifest
progTransport *helpers.Transport
targetPlatform *ocispec.Platform
log *slog.Logger
repo *remote.Repository
cache *oci.Store
root *Manifest
progTransport *helpers.Transport
targetPlatform *ocispec.Platform
insecureSkipVerify *bool
log *slog.Logger
}

// Modifier is a function that modifies an OrasRemote
Expand All @@ -46,12 +47,16 @@ func WithPlainHTTP(plainHTTP bool) Modifier {
}
}

// WithInsecureSkipVerify sets the insecure TLS flag for the remote
// WithInsecureSkipVerify sets the insecure TLS flag for the remote.
// An explicit value takes precedence over WithTransport regardless of modifier order.
func WithInsecureSkipVerify(insecure bool) Modifier {
return func(o *OrasRemote) {
o.insecureSkipVerify = &insecure
transport, ok := o.progTransport.Base.(*http.Transport)
if ok {
transport.TLSClientConfig.InsecureSkipVerify = insecure
transport = transport.Clone()
applyInsecureSkipVerify(transport, insecure)
o.progTransport.Base = transport
return
}
if o.log != nil {
Expand All @@ -60,6 +65,19 @@ func WithInsecureSkipVerify(insecure bool) Modifier {
}
}

// WithTransport sets the HTTP transport for the remote.
func WithTransport(transport *http.Transport) Modifier {
Comment thread
Racer159 marked this conversation as resolved.
return func(o *OrasRemote) {
if transport != nil {
transport = transport.Clone()
if o.insecureSkipVerify != nil {
applyInsecureSkipVerify(transport, *o.insecureSkipVerify)
}
o.progTransport.Base = transport
}
}
}

// PlatformForArch sets the target architecture for the remote
func PlatformForArch(arch string) ocispec.Platform {
return ocispec.Platform{
Expand Down Expand Up @@ -109,17 +127,17 @@ func NewOrasRemote(url string, platform ocispec.Platform, mods ...Modifier) (*Or
return nil, fmt.Errorf("http.DefaultTransport is not an *http.Transport, something mutated global net/http variables")
}
transport := httpTransport.Clone()
progTransport := helpers.NewTransport(transport, nil)
client := &auth.Client{
Client: retry.DefaultClient,
Client: &http.Client{Transport: progTransport},
Header: http.Header{
"User-Agent": {"oras-go"},
},
Cache: auth.DefaultCache,
Cache: auth.NewCache(),
}
client.Client.Transport = transport
o := &OrasRemote{
repo: &remote.Repository{Client: client},
progTransport: helpers.NewTransport(transport, nil),
progTransport: progTransport,
targetPlatform: &platform,
log: slog.Default(),
}
Expand Down Expand Up @@ -189,15 +207,24 @@ func (o *OrasRemote) setRepository(ref registry.Reference) error {
if err != nil {
return fmt.Errorf("failed to get credentials: %w", err)
}
client := &auth.Client{
Client: retry.DefaultClient,
Cache: auth.NewCache(),
Credential: credentials.Credential(credStore),
client, ok := o.repo.Client.(*auth.Client)
if !ok {
return fmt.Errorf("repository client is not an auth client")
}
client.Credential = credentials.Credential(credStore)
if o.log != nil {
o.log.Debug("gathering credentials from default Docker config file", "credentials_configured", credStore.IsAuthConfigured())
}
o.log.Debug("gathering credentials from default Docker config file", "credentials_configured", credStore.IsAuthConfigured())

o.repo.Reference = ref
o.repo.Client = client

return nil
}

// applyInsecureSkipVerify sets InsecureSkipVerify on the TLSClientConfig
func applyInsecureSkipVerify(transport *http.Transport, insecure bool) {
if transport.TLSClientConfig == nil {
transport.TLSClientConfig = &tls.Config{}
}
transport.TLSClientConfig.InsecureSkipVerify = insecure
}
203 changes: 203 additions & 0 deletions oci/common_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,203 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: 2024-Present Defense Unicorns

package oci

import (
"crypto/tls"
"net/http"
"net/http/httptest"
"testing"

"github.com/stretchr/testify/require"
"oras.land/oras-go/v2/registry/remote/auth"
"oras.land/oras-go/v2/registry/remote/retry"
)

func TestNewOrasRemote_TransportIsolation(t *testing.T) {
retryTransport := retry.DefaultClient.Transport
platform := PlatformForArch(testArch)

first, err := NewOrasRemote("example.com/first:latest", platform)
require.NoError(t, err)
second, err := NewOrasRemote("example.com/second:latest", platform)
require.NoError(t, err)

firstClient, ok := first.repo.Client.(*auth.Client)
require.True(t, ok)
secondClient, ok := second.repo.Client.(*auth.Client)
require.True(t, ok)
require.NotSame(t, firstClient.Client, secondClient.Client)
require.NotSame(t, firstClient.Cache, secondClient.Cache)
require.NotSame(t, first.progTransport, second.progTransport)
require.NotSame(t, first.progTransport.Base, second.progTransport.Base)
require.Same(t, retryTransport, retry.DefaultClient.Transport)
}

func TestWithTransport_SurvivesProgressChanges(t *testing.T) {
transport := &http.Transport{}
remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
WithTransport(transport),
)
require.NoError(t, err)
require.NotSame(t, transport, remote.progTransport.Base)
configuredTransport, ok := remote.progTransport.Base.(*http.Transport)
require.True(t, ok)

client, ok := remote.repo.Client.(*auth.Client)
require.True(t, ok)
require.Same(t, remote.progTransport, client.Client.Transport)

remote.SetProgressWriter(&TestProgressWriter{})
require.Same(t, configuredTransport, remote.progTransport.Base)
require.Same(t, remote.progTransport, client.Client.Transport)

remote.ClearProgressWriter()
require.Same(t, configuredTransport, remote.progTransport.Base)
require.Same(t, remote.progTransport, client.Client.Transport)
}

func TestWithTransport_UsedForRequests(t *testing.T) {
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)

transport, ok := server.Client().Transport.(*http.Transport)
require.True(t, ok)
remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
WithTransport(transport),
)
require.NoError(t, err)

client, ok := remote.repo.Client.(*auth.Client)
require.True(t, ok)
response, err := client.Client.Get(server.URL)
require.NoError(t, err)
require.NoError(t, response.Body.Close())
require.Equal(t, http.StatusNoContent, response.StatusCode)
}

func TestWithUserAgent_SurvivesRepositorySetup(t *testing.T) {
remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
WithUserAgent("zarf/test"),
)
require.NoError(t, err)

client, ok := remote.repo.Client.(*auth.Client)
require.True(t, ok)
require.Equal(t, "zarf/test", client.Header.Get("User-Agent"))
require.NotNil(t, client.Credential)
}

func TestWithLogger_Nil(t *testing.T) {
remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
WithLogger(nil),
)
require.NoError(t, err)
require.Nil(t, remote.log)
}

func TestWithInsecureSkipVerify_ClonesTransport(t *testing.T) {
tests := []struct {
name string
tlsConfig *tls.Config
}{
{name: "nil TLS config"},
{name: "existing TLS config", tlsConfig: &tls.Config{ServerName: "registry.example.com"}},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
transport := &http.Transport{TLSClientConfig: tt.tlsConfig}
remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
WithTransport(transport),
WithInsecureSkipVerify(true),
)
require.NoError(t, err)

configured, ok := remote.progTransport.Base.(*http.Transport)
require.True(t, ok)
require.NotSame(t, transport, configured)
require.NotNil(t, configured.TLSClientConfig)
require.True(t, configured.TLSClientConfig.InsecureSkipVerify)
require.False(t, transport.TLSClientConfig.InsecureSkipVerify)
if tt.tlsConfig != nil {
require.NotSame(t, tt.tlsConfig, configured.TLSClientConfig)
require.Equal(t, tt.tlsConfig.ServerName, configured.TLSClientConfig.ServerName)
require.False(t, tt.tlsConfig.InsecureSkipVerify)
}
})
}
}

func TestWithInsecureSkipVerify_OrderIndependent(t *testing.T) {
tests := []struct {
name string
insecure bool
insecureFirst bool
}{
{name: "enable after transport", insecure: true, insecureFirst: false},
{name: "enable before transport", insecure: true, insecureFirst: true},
{name: "disable after transport", insecure: false, insecureFirst: false},
{name: "disable before transport", insecure: false, insecureFirst: true},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
transport := &http.Transport{TLSClientConfig: &tls.Config{InsecureSkipVerify: !tt.insecure}}
mods := []Modifier{WithTransport(transport), WithInsecureSkipVerify(tt.insecure)}
if tt.insecureFirst {
mods[0], mods[1] = mods[1], mods[0]
}

remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
mods...,
)
require.NoError(t, err)

configured, ok := remote.progTransport.Base.(*http.Transport)
require.True(t, ok)
require.Equal(t, tt.insecure, configured.TLSClientConfig.InsecureSkipVerify)
require.Equal(t, !tt.insecure, transport.TLSClientConfig.InsecureSkipVerify)
})
}
}

func TestWithTransport_PreservesInsecureSkipVerifyWhenUnset(t *testing.T) {
tests := []struct {
name string
insecure bool
}{
{name: "secure", insecure: false},
{name: "insecure", insecure: true},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
transport := &http.Transport{TLSClientConfig: &tls.Config{InsecureSkipVerify: tt.insecure}}
remote, err := NewOrasRemote(
"example.com/repository:latest",
PlatformForArch(testArch),
WithTransport(transport),
)
require.NoError(t, err)

configured, ok := remote.progTransport.Base.(*http.Transport)
require.True(t, ok)
require.Equal(t, tt.insecure, configured.TLSClientConfig.InsecureSkipVerify)
})
}
}
Loading