diff --git a/cmd/nvfleetint/node_test.go b/cmd/nvfleetint/node_test.go index 3804694..bf79549 100644 --- a/cmd/nvfleetint/node_test.go +++ b/cmd/nvfleetint/node_test.go @@ -10,6 +10,7 @@ import ( "net/http/httptest" "slices" "strings" + "sync/atomic" "testing" "github.com/NVIDIA/fleet-intelligence-client/nvfleetint" @@ -287,6 +288,64 @@ func TestNodeListAllDefaultsPageSize(t *testing.T) { } } +// Verifies --all retries only a transiently failing page instead of restarting +// pagination and duplicating already collected nodes. +func TestNodeListAllRetriesFailedPage(t *testing.T) { + t.Setenv("HOME", t.TempDir()) + + var firstPageCalls atomic.Int32 + var secondPageCalls atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Query().Get("page") { + case "0": + firstPageCalls.Add(1) + _, _ = w.Write([]byte(`{"nodes":[{"nodeUUID":"node-1"}],"hasMore":true,"page":0,"pageSize":100,"total":2}`)) + case "1": + if secondPageCalls.Add(1) == 1 { + w.WriteHeader(http.StatusServiceUnavailable) + _, _ = w.Write([]byte(`{"error":"temporarily unavailable"}`)) + return + } + _, _ = w.Write([]byte(`{"nodes":[{"nodeUUID":"node-2"}],"hasMore":false,"page":1,"pageSize":100,"total":2}`)) + default: + t.Fatalf("unexpected page: %q", r.URL.Query().Get("page")) + } + })) + defer server.Close() + + saveTestConfig(t, server.URL, "test-key") + var out bytes.Buffer + cmd := newRootCmd() + cmd.SetOut(&out) + cmd.SetErr(&bytes.Buffer{}) + cmd.SetArgs([]string{"node", "list", "--all", "--output", "json", "--timeout", "5s"}) + if err := cmd.Execute(); err != nil { + t.Fatalf("execute failed: %v", err) + } + + var got struct { + Items []struct { + NodeUUID string `json:"nodeUUID"` + } `json:"items"` + Pagination struct { + PagesFetched int `json:"pagesFetched"` + } `json:"pagination"` + } + if err := json.Unmarshal(out.Bytes(), &got); err != nil { + t.Fatalf("decode output: %v", err) + } + if len(got.Items) != 2 || + got.Items[0].NodeUUID != "node-1" || + got.Items[1].NodeUUID != "node-2" || + got.Pagination.PagesFetched != 2 { + t.Fatalf("unexpected merged output: %#v", got) + } + if firstPageCalls.Load() != 1 || secondPageCalls.Load() != 2 { + t.Fatalf("unexpected page calls: first=%d second=%d", firstPageCalls.Load(), secondPageCalls.Load()) + } +} + // Verifies node describe table output func TestNodeDescribeTable(t *testing.T) { t.Setenv("HOME", t.TempDir()) diff --git a/internal/clihelpers/pagination.go b/internal/clihelpers/pagination.go index 207f6a5..7859a2a 100644 --- a/internal/clihelpers/pagination.go +++ b/internal/clihelpers/pagination.go @@ -53,7 +53,10 @@ func FetchAllRawPages(itemKey string, startPage int, fetch func(page int) (RawPa page, err := fetch(startPage + offset) if err != nil { - return MergedJSONResult{}, err + return MergedJSONResult{}, fmt.Errorf( + "fetch %s API page %d after %d completed pages: %w", + itemKey, startPage+offset, result.Pagination.PagesFetched, err, + ) } items, err := ExtractRawItems(page.RawJSON, itemKey) if err != nil { diff --git a/nvfleetint/client.go b/nvfleetint/client.go index d03d25e..b0afb28 100644 --- a/nvfleetint/client.go +++ b/nvfleetint/client.go @@ -9,11 +9,14 @@ import ( "errors" "fmt" "io" + "math/rand/v2" "net" "net/http" "net/url" + "strconv" "strings" "sync" + "syscall" "time" "github.com/NVIDIA/fleet-intelligence-client/internal/generated/fleetapi" @@ -22,6 +25,13 @@ import ( // DefaultTimeout is the per-request timeout applied when none is configured. const DefaultTimeout = 2 * time.Minute +const ( + defaultRequestAttempts = 3 + initialRetryDelay = 200 * time.Millisecond + maximumRetryDelay = 5 * time.Second + maximumRetryAfterSecs = int64(1<<63-1) / int64(time.Second) +) + // signingKeyPath is the well-known location of the report signing public key. const signingKeyPath = "/.well-known/signing-key.pub" @@ -36,11 +46,12 @@ var ( // Calls the Fleet Intelligence customer API type Client struct { - baseURL *url.URL - apiKey string - httpClient *http.Client - timeout time.Duration - api *fleetapi.ClientWithResponses + baseURL *url.URL + apiKey string + httpClient *http.Client + requestDoer fleetapi.HttpRequestDoer + timeout time.Duration + api *fleetapi.ClientWithResponses } // Customizes client construction behavior @@ -104,10 +115,17 @@ func NewClient(baseURL, apiKey string, opts ...Option) (*Client, error) { for _, opt := range opts { opt(client) } + client.requestDoer = &retryingDoer{ + inner: client.httpClient, + maxAttempts: defaultRequestAttempts, + } api, err := fleetapi.NewClientWithResponses( client.baseURL.String(), - fleetapi.WithHTTPClient(&timeoutDoer{inner: client.httpClient, timeout: client.timeout}), + fleetapi.WithHTTPClient(&timeoutDoer{ + inner: client.requestDoer, + timeout: client.timeout, + }), fleetapi.WithRequestEditorFn(client.authorizeRequest), ) if err != nil { @@ -248,6 +266,156 @@ type timeoutDoer struct { timeout time.Duration } +// Retries idempotent reads when a transient transport or backend failure makes +// an individual request fail. Paginated callers therefore retry only the page +// that failed; successfully returned pages are not requested again. +type retryingDoer struct { + inner fleetapi.HttpRequestDoer + maxAttempts int + delay func(attempt int, response *http.Response) time.Duration + wait func(context.Context, time.Duration) error +} + +// Performs one request, retrying bounded transient failures within the request +// context's existing timeout. +func (d *retryingDoer) Do(req *http.Request) (*http.Response, error) { + maxAttempts := d.maxAttempts + if maxAttempts < 1 { + maxAttempts = 1 + } + + for attempt := 1; attempt <= maxAttempts; attempt++ { + response, err := d.inner.Do(req) + if !canRetryRequest(req, response, err) || attempt == maxAttempts { + return response, err + } + + if response != nil && response.Body != nil { + _, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 4<<10)) + _ = response.Body.Close() + } + + delay := defaultRetryDelay(attempt, response) + if d.delay != nil { + delay = d.delay(attempt, response) + } + wait := waitForRetry + if d.wait != nil { + wait = d.wait + } + if err := wait(req.Context(), delay); err != nil { + return nil, err + } + } + + return nil, errors.New("retry attempts exhausted") +} + +// Reports whether retrying this request is safe and potentially useful. +func canRetryRequest(req *http.Request, response *http.Response, err error) bool { + if req == nil || (req.Method != http.MethodGet && req.Method != http.MethodHead) { + return false + } + // The generated read requests have no body. Avoid replaying an unusual GET + // body because net/http may already have consumed or closed it. + if req.Body != nil { + return false + } + if req.Context().Err() != nil { + return false + } + if err != nil { + return isRetryableNetworkError(err) + } + if response == nil { + return false + } + + switch response.StatusCode { + case http.StatusRequestTimeout, + http.StatusTooManyRequests, + http.StatusInternalServerError, + http.StatusBadGateway, + http.StatusServiceUnavailable, + http.StatusGatewayTimeout: + return true + default: + return false + } +} + +// Identifies transport failures that commonly succeed when retried. +func isRetryableNetworkError(err error) bool { + if err == nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return false + } + + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + + return errors.Is(err, io.EOF) || + errors.Is(err, io.ErrUnexpectedEOF) || + errors.Is(err, syscall.ECONNRESET) || + errors.Is(err, syscall.ECONNREFUSED) || + errors.Is(err, syscall.EPIPE) +} + +// Computes Retry-After or an exponential delay with bounded jitter. +func defaultRetryDelay(attempt int, response *http.Response) time.Duration { + if delay, ok := responseRetryAfter(response, time.Now()); ok { + return delay + } + + delay := initialRetryDelay << (attempt - 1) + if delay > maximumRetryDelay { + delay = maximumRetryDelay + } + // Randomize to 50%-150% so many clients do not retry in lockstep. + half := delay / 2 + return half + time.Duration(rand.Int64N(int64(delay)+1)) +} + +// Parses Retry-After as either seconds or an HTTP date. +func responseRetryAfter(response *http.Response, now time.Time) (time.Duration, bool) { + if response == nil { + return 0, false + } + raw := strings.TrimSpace(response.Header.Get("Retry-After")) + if raw == "" { + return 0, false + } + if seconds, err := strconv.ParseInt(raw, 10, 64); err == nil && + seconds >= 0 && seconds <= maximumRetryAfterSecs { + return time.Duration(seconds) * time.Second, true + } + retryAt, err := http.ParseTime(raw) + if err != nil { + return 0, false + } + if !retryAt.After(now) { + return 0, true + } + return retryAt.Sub(now), true +} + +// Waits for a retry delay without ignoring request cancellation. +func waitForRetry(ctx context.Context, delay time.Duration) error { + if delay <= 0 { + return nil + } + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-timer.C: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + // Performs the request and rewrites timeout errors into a friendly message func (d *timeoutDoer) Do(req *http.Request) (*http.Response, error) { resp, err := d.inner.Do(req) @@ -293,7 +461,7 @@ func (c *Client) FetchSigningKey(ctx context.Context) ([]byte, error) { } req.Header.Set("Accept", signingKeyAcceptHeader) - resp, err := c.httpClient.Do(req) + resp, err := c.requestDoer.Do(req) if err != nil { if errors.Is(err, context.DeadlineExceeded) { return nil, fmt.Errorf("request timed out after %s", c.timeout) diff --git a/nvfleetint/client_test.go b/nvfleetint/client_test.go index a695d3f..c75912b 100644 --- a/nvfleetint/client_test.go +++ b/nvfleetint/client_test.go @@ -4,13 +4,17 @@ package nvfleetint import ( + "bytes" "context" "crypto/tls" "errors" + "io" "net/http" "net/http/httptest" "net/url" + "strconv" "strings" + "sync/atomic" "testing" "time" ) @@ -97,6 +101,7 @@ func TestHardenedTransportPreservesStricterTLSMinVersion(t *testing.T) { if hardened == nil { t.Fatal("expected a hardened transport") + return } if hardened.TLSClientConfig.MinVersion != tls.VersionTLS13 { t.Fatalf("expected TLS 1.3 MinVersion to be preserved, got %d", hardened.TLSClientConfig.MinVersion) @@ -109,9 +114,11 @@ func TestHardenedTransportAddsMissingTLSConfig(t *testing.T) { if hardened == nil { t.Fatal("expected a hardened transport") + return } if hardened.TLSClientConfig == nil { t.Fatal("expected a TLS config to be added") + return } if hardened.TLSClientConfig.MinVersion != tls.VersionTLS12 { t.Fatalf("unexpected TLS MinVersion: %d", hardened.TLSClientConfig.MinVersion) @@ -132,6 +139,130 @@ func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { return f(req) } +// Verifies transient HTTP failures retry the same idempotent request and return +// the first successful response. +func TestRetryingDoerRetriesTransientStatus(t *testing.T) { + var calls atomic.Int32 + inner := &http.Client{Transport: roundTripperFunc(func(req *http.Request) (*http.Response, error) { + call := calls.Add(1) + status := http.StatusServiceUnavailable + if call == defaultRequestAttempts { + status = http.StatusOK + } + return &http.Response{ + StatusCode: status, + Header: make(http.Header), + Body: io.NopCloser(bytes.NewBufferString(http.StatusText(status))), + Request: req, + }, nil + })} + doer := &retryingDoer{ + inner: inner, + maxAttempts: defaultRequestAttempts, + delay: func(int, *http.Response) time.Duration { return 0 }, + } + + req, err := http.NewRequest(http.MethodGet, "https://example.com/v1/nodes?page=4", nil) + if err != nil { + t.Fatalf("build request: %v", err) + } + response, err := doer.Do(req) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("unexpected status: %d", response.StatusCode) + } + if got := calls.Load(); got != defaultRequestAttempts { + t.Fatalf("expected %d attempts, got %d", defaultRequestAttempts, got) + } +} + +// Verifies permanent HTTP failures are returned without retrying. +func TestRetryingDoerDoesNotRetryPermanentStatus(t *testing.T) { + var calls atomic.Int32 + inner := &http.Client{Transport: roundTripperFunc(func(req *http.Request) (*http.Response, error) { + calls.Add(1) + return &http.Response{ + StatusCode: http.StatusBadRequest, + Header: make(http.Header), + Body: io.NopCloser(bytes.NewBufferString("bad request")), + Request: req, + }, nil + })} + doer := &retryingDoer{inner: inner, maxAttempts: defaultRequestAttempts} + + req, err := http.NewRequest(http.MethodGet, "https://example.com/v1/nodes", nil) + if err != nil { + t.Fatalf("build request: %v", err) + } + response, err := doer.Do(req) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusBadRequest || calls.Load() != 1 { + t.Fatalf("unexpected response status/calls: %d/%d", response.StatusCode, calls.Load()) + } +} + +// Verifies Retry-After supports both seconds and HTTP-date forms. +func TestResponseRetryAfter(t *testing.T) { + now := time.Date(2026, time.August, 5, 12, 0, 0, 0, time.UTC) + maximumDelay := time.Duration(maximumRetryAfterSecs) * time.Second + tests := []struct { + name string + raw string + want time.Duration + ok bool + }{ + {name: "seconds", raw: "7", want: 7 * time.Second, ok: true}, + {name: "maximum seconds", raw: strconv.FormatInt(maximumRetryAfterSecs, 10), want: maximumDelay, ok: true}, + {name: "overflowing seconds", raw: strconv.FormatInt(maximumRetryAfterSecs+1, 10), ok: false}, + {name: "date", raw: now.Add(11 * time.Second).Format(http.TimeFormat), want: 11 * time.Second, ok: true}, + {name: "past date", raw: now.Add(-time.Second).Format(http.TimeFormat), want: 0, ok: true}, + {name: "invalid", raw: "later", ok: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + response := &http.Response{Header: http.Header{"Retry-After": []string{tt.raw}}} + got, ok := responseRetryAfter(response, now) + if ok != tt.ok || got != tt.want { + t.Fatalf("unexpected retry delay: got %v/%t want %v/%t", got, ok, tt.want, tt.ok) + } + }) + } +} + +// Verifies an oversized Retry-After value falls back to a positive bounded +// delay, so the retry wait cannot be bypassed by duration overflow. +func TestOversizedRetryAfterUsesBackoff(t *testing.T) { + response := &http.Response{Header: http.Header{ + "Retry-After": []string{strconv.FormatInt(maximumRetryAfterSecs+1, 10)}, + }} + delay := defaultRetryDelay(1, response) + if delay <= 0 { + t.Fatalf("expected positive fallback delay, got %v", delay) + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := waitForRetry(ctx, delay); !errors.Is(err, context.Canceled) { + t.Fatalf("expected retry wait to honor cancellation, got %v", err) + } +} + +// Verifies a canceled request interrupts retry backoff. +func TestWaitForRetryHonorsCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := waitForRetry(ctx, time.Minute); !errors.Is(err, context.Canceled) { + t.Fatalf("expected context cancellation, got %v", err) + } +} + // Verifies client configuration accessors func TestNewClientStoresConfiguration(t *testing.T) { client, err := NewClient("https://example.com", "key") diff --git a/nvfleetint/verify_test.go b/nvfleetint/verify_test.go index 16566f4..38c7616 100644 --- a/nvfleetint/verify_test.go +++ b/nvfleetint/verify_test.go @@ -8,7 +8,9 @@ import ( "net/http" "net/http/httptest" "strings" + "sync/atomic" "testing" + "time" "github.com/sigstore/sigstore-go/pkg/bundle" "github.com/sigstore/sigstore-go/pkg/sign" @@ -130,6 +132,38 @@ func TestFetchSigningKey(t *testing.T) { } } +// Verifies direct SDK requests use the same central retry behavior as generated +// API calls. +func TestFetchSigningKeyRetriesTransientFailure(t *testing.T) { + var calls atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if calls.Add(1) == 1 { + http.Error(w, "temporarily unavailable", http.StatusServiceUnavailable) + return + } + _, _ = w.Write([]byte("public key")) + })) + defer server.Close() + + client, err := NewClient(server.URL, "test-key") + if err != nil { + t.Fatalf("new client failed: %v", err) + } + retryer, ok := client.requestDoer.(*retryingDoer) + if !ok { + t.Fatalf("unexpected request doer: %T", client.requestDoer) + } + retryer.delay = func(int, *http.Response) time.Duration { return 0 } + + key, err := client.FetchSigningKey(context.Background()) + if err != nil { + t.Fatalf("fetch signing key failed: %v", err) + } + if string(key) != "public key" || calls.Load() != 2 { + t.Fatalf("unexpected key/calls: %q/%d", key, calls.Load()) + } +} + // Verifies a non-200 response from the key endpoint is surfaced as an error func TestFetchSigningKeyError(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {