Skip to content
Open
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
28 changes: 22 additions & 6 deletions pkg/files/reader.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ import (
"strings"
"time"

"github.com/aws/eks-anywhere/pkg/retrier"
"golang.org/x/net/http/httpproxy"
)

Expand All @@ -24,6 +25,7 @@ type Reader struct {
embedFS embed.FS
httpClient *http.Client
userAgent string
retrier *retrier.Retrier
}

type ReaderOpt func(*Reader)
Expand All @@ -47,6 +49,14 @@ func WithEKSAUserAgent(eksAComponent, version string) ReaderOpt {
return WithUserAgent(eksaUserAgent(eksAComponent, version))
}

// WithRetrier allows to use a custom retrier for the http GET requests
// performed when reading a file from a url. This is only for testing.
func WithRetrier(retrier *retrier.Retrier) ReaderOpt {
return func(r *Reader) {
r.retrier = retrier
}
}

// WithRootCACerts configures the HTTP client's trusted CAs. Note that this will overwrite
// the defaults so the host's trust will be ignored. This option is only for testing.
func WithRootCACerts(certs []*x509.Certificate) ReaderOpt {
Expand Down Expand Up @@ -99,6 +109,7 @@ func NewReader(opts ...ReaderOpt) *Reader {
embedFS: embedFS,
httpClient: client,
userAgent: eksaUserAgent("unknown", "no-version"),
retrier: retrier.NewWithMaxRetries(5, 5*time.Second),
}

for _, o := range opts {
Expand Down Expand Up @@ -131,13 +142,18 @@ func (r *Reader) readHttpFile(uri string) ([]byte, error) {
}

request.Header.Set("User-Agent", r.userAgent)
resp, err := r.httpClient.Do(request)
if err != nil {
return nil, fmt.Errorf("failed reading file from url [%s]: %v", uri, err)
}
defer resp.Body.Close()

data, err := io.ReadAll(resp.Body)
var data []byte
err = r.retrier.Retry(func() error {
resp, err := r.httpClient.Do(request)
if err != nil {
return err
}
defer resp.Body.Close()

data, err = io.ReadAll(resp.Body)
return err
})
if err != nil {
return nil, fmt.Errorf("failed reading file from url [%s]: %v", uri, err)
}
Expand Down
65 changes: 65 additions & 0 deletions pkg/files/reader_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,13 +11,15 @@ import (
"net/url"
"os"
"sync"
"sync/atomic"
"testing"
"time"

. "github.com/onsi/gomega"

"github.com/aws/eks-anywhere/internal/test"
"github.com/aws/eks-anywhere/pkg/files"
"github.com/aws/eks-anywhere/pkg/retrier"
)

//go:embed testdata
Expand Down Expand Up @@ -95,6 +97,69 @@ func TestReaderReadFileHTTPSSuccess(t *testing.T) {
test.AssertContentToFile(t, string(got), filePath)
}

func TestReaderReadFileHTTPSRetriesOnTransientError(t *testing.T) {
g := NewWithT(t)
filePath := "testdata/file.yaml"

var attempts int32
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if atomic.AddInt32(&attempts, 1) <= 2 {
// Simulate a transient network error by closing the connection
// before a response is sent back to the client.
hj, ok := w.(http.Hijacker)
g.Expect(ok).To(BeTrue())
conn, _, err := hj.Hijack()
g.Expect(err).To(BeNil())
conn.Close()
return
}

fileContent, err := os.ReadFile(filePath)
g.Expect(err).To(BeNil())
if _, err := w.Write(fileContent); err != nil {
t.Errorf("Failed writing response to http request: %s", err)
}
}))
t.Cleanup(func() { server.Close() })

uri := server.URL + "/" + filePath

r := files.NewReader(
files.WithRootCACerts(serverCerts(g, server)),
files.WithRetrier(retrier.NewWithMaxRetries(5, 0)),
)
got, err := r.ReadFile(uri)
g.Expect(err).To(BeNil())
test.AssertContentToFile(t, string(got), filePath)
g.Expect(atomic.LoadInt32(&attempts)).To(Equal(int32(3)))
}

func TestReaderReadFileHTTPSFailsAfterExhaustingRetries(t *testing.T) {
g := NewWithT(t)
filePath := "testdata/file.yaml"

var attempts int32
server := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
atomic.AddInt32(&attempts, 1)
hj, ok := w.(http.Hijacker)
g.Expect(ok).To(BeTrue())
conn, _, err := hj.Hijack()
g.Expect(err).To(BeNil())
conn.Close()
}))
t.Cleanup(func() { server.Close() })

uri := server.URL + "/" + filePath

r := files.NewReader(
files.WithRootCACerts(serverCerts(g, server)),
files.WithRetrier(retrier.NewWithMaxRetries(3, 0)),
)
_, err := r.ReadFile(uri)
g.Expect(err).NotTo(BeNil())
g.Expect(atomic.LoadInt32(&attempts)).To(Equal(int32(3)))
}

func TestReaderReadFileHTTPSProxySuccess(t *testing.T) {
t.Skip("Flaky (https://github.com/aws/eks-anywhere/issues/5775)")

Expand Down