diff --git a/CHANGELOG.md b/CHANGELOG.md index 0f6b6930..3824360a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,9 @@ ### Fixed +- Retry rate-limited feature flag requests using applicable quota reset times and `Retry-After` +- Preserve cancellation and timeout errors while reading feature flag responses + ## 2.3.1 ### Added diff --git a/go.mod b/go.mod index aa02e332..59394ba7 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,7 @@ go 1.25.0 require ( github.com/Masterminds/sprig/v3 v3.3.0 github.com/google/go-querystring v1.1.0 // indirect + github.com/hashicorp/go-retryablehttp v0.7.8 github.com/monochromegane/go-gitignore v0.0.0-20200626010858-205db1a8cc00 github.com/sourcegraph/go-diff v0.6.1 github.com/spf13/viper v1.21.0 @@ -50,7 +51,6 @@ require ( github.com/go-git/go-git/v5 v5.19.1 // indirect github.com/google/uuid v1.6.0 // indirect github.com/hashicorp/go-cleanhttp v0.5.2 // indirect - github.com/hashicorp/go-retryablehttp v0.7.8 // indirect github.com/huandu/xstrings v1.5.0 // indirect github.com/iancoleman/strcase v0.3.0 // indirect github.com/kyoh86/xdg v1.2.0 // indirect diff --git a/internal/ldclient/diagnostics.go b/internal/ldclient/diagnostics.go new file mode 100644 index 00000000..7e8384df --- /dev/null +++ b/internal/ldclient/diagnostics.go @@ -0,0 +1,205 @@ +package ldapi + +import ( + "encoding/json" + "time" +) + +const diagnosticsSchema = 1 + +type diagnosticLog func(string, ...any) + +type fetchDiagnostics struct { + log diagnosticLog + startedAt time.Time + + httpAttempts int + retries int + responses429 int + successfulPages int + nonemptyPages int + activeFlags int + archivedFlags int + retryWaits int + proactiveWaits int + scheduledWait time.Duration + lastHTTPStatus *int + + failedCollection *string + failedOffset *int +} + +func newFetchDiagnostics(startedAt time.Time, log diagnosticLog) *fetchDiagnostics { + if log == nil { + log = func(string, ...any) {} + } + return &fetchDiagnostics{log: log, startedAt: startedAt} +} + +func (d *fetchDiagnostics) emitStart() { + d.emit(struct { + Schema int `json:"schema"` + Event string `json:"event"` + }{ + Schema: diagnosticsSchema, + Event: "start", + }) +} + +type rateLimitDiagnostic struct { + Schema int `json:"schema"` + Event string `json:"event"` + Collection string `json:"collection"` + Offset int `json:"offset"` + Attempt int `json:"attempt"` + Status int `json:"status"` + GlobalResetMillis *int64 `json:"global_reset_ms"` + AuthTokenResetMillis *int64 `json:"auth_token_reset_ms"` + RetryAfterSeconds *int64 `json:"retry_after_seconds"` + RetryAfterDateMillis *int64 `json:"retry_after_date_ms"` + GlobalRemaining *int64 `json:"global_remaining"` + RouteRemaining *int64 `json:"route_remaining"` + AuthTokenRemaining *int64 `json:"auth_token_remaining"` +} + +func (d *fetchDiagnostics) emitRateLimit(collection string, offset, attempt, status int, hints parsedRateLimitHeaders) { + d.emit(rateLimitDiagnostic{ + Schema: diagnosticsSchema, + Event: "rate_limit", + Collection: collection, + Offset: offset, + Attempt: attempt, + Status: status, + GlobalResetMillis: hints.globalResetMillis(), + AuthTokenResetMillis: hints.authTokenResetMillis(), + RetryAfterSeconds: hints.retryAfterSecondsValue(), + RetryAfterDateMillis: hints.retryAfterDateMillis(), + GlobalRemaining: hints.globalRemainingValue(), + RouteRemaining: hints.routeRemainingValue(), + AuthTokenRemaining: hints.authTokenRemainingValue(), + }) +} + +type waitDiagnostic struct { + Schema int `json:"schema"` + Event string `json:"event"` + Collection string `json:"collection"` + Offset int `json:"offset"` + NextAttempt int `json:"next_attempt"` + Reason string `json:"reason"` + IntendedWaitMillis int64 `json:"intended_wait_ms"` +} + +func (d *fetchDiagnostics) emitWait(collection string, offset, nextAttempt int, reason string, wait time.Duration) { + d.emit(waitDiagnostic{ + Schema: diagnosticsSchema, + Event: "wait", + Collection: collection, + Offset: offset, + NextAttempt: nextAttempt, + Reason: reason, + IntendedWaitMillis: wait.Milliseconds(), + }) +} + +type summaryDiagnostic struct { + Schema int `json:"schema"` + Event string `json:"event"` + Outcome string `json:"outcome"` + InventoryComplete bool `json:"inventory_complete"` + HTTPAttempts int `json:"http_attempts"` + Retries int `json:"retries"` + Responses429 int `json:"responses_429"` + SuccessfulPages int `json:"successful_pages"` + NonemptyPages int `json:"nonempty_pages"` + ActiveFlags int `json:"active_flags"` + ArchivedFlags int `json:"archived_flags"` + RetryWaits int `json:"retry_waits"` + ProactiveWaits int `json:"proactive_waits"` + ScheduledWaitMs int64 `json:"scheduled_wait_ms"` + ElapsedMs int64 `json:"elapsed_ms"` + LastHTTPStatus *int `json:"last_http_status"` + FailedCollection *string `json:"failed_collection"` + FailedOffset *int `json:"failed_offset"` +} + +func (d *fetchDiagnostics) emitSummary(now time.Time, outcome string, complete bool) { + elapsed := now.Sub(d.startedAt) + if elapsed < 0 { + elapsed = 0 + } + d.emit(summaryDiagnostic{ + Schema: diagnosticsSchema, + Event: "summary", + Outcome: outcome, + InventoryComplete: complete, + HTTPAttempts: d.httpAttempts, + Retries: d.retries, + Responses429: d.responses429, + SuccessfulPages: d.successfulPages, + NonemptyPages: d.nonemptyPages, + ActiveFlags: d.activeFlags, + ArchivedFlags: d.archivedFlags, + RetryWaits: d.retryWaits, + ProactiveWaits: d.proactiveWaits, + ScheduledWaitMs: d.scheduledWait.Milliseconds(), + ElapsedMs: elapsed.Milliseconds(), + LastHTTPStatus: d.lastHTTPStatus, + FailedCollection: d.failedCollection, + FailedOffset: d.failedOffset, + }) +} + +func (d *fetchDiagnostics) emit(value any) { + encoded, err := json.Marshal(value) + if err != nil { + return + } + d.log("LD_FLAG_FETCH %s\n", encoded) +} + +func (d *fetchDiagnostics) recordAttempt(retryNumber int) { + d.httpAttempts++ + if retryNumber > 0 { + d.retries++ + } +} + +func (d *fetchDiagnostics) recordStatus(status int) { + d.lastHTTPStatus = intPointer(status) +} + +func (d *fetchDiagnostics) recordPage(collection string, itemCount int) { + d.successfulPages++ + if itemCount > 0 { + d.nonemptyPages++ + } + if collection == flagCollectionArchived { + d.archivedFlags += itemCount + } else { + d.activeFlags += itemCount + } +} + +func (d *fetchDiagnostics) recordFailure(collection string, offset int) { + d.failedCollection = stringPointer(collection) + d.failedOffset = intPointer(offset) +} + +func (d *fetchDiagnostics) recordRetryWait(wait time.Duration) { + d.retryWaits++ + d.scheduledWait += wait +} + +func (d *fetchDiagnostics) recordProactiveWait(wait time.Duration) { + d.proactiveWaits++ + d.scheduledWait += wait +} + +func intPointer(value int) *int { + return &value +} + +func stringPointer(value string) *string { + return &value +} diff --git a/internal/ldclient/flags.go b/internal/ldclient/flags.go index 4ea77561..f8ed807e 100644 --- a/internal/ldclient/flags.go +++ b/internal/ldclient/flags.go @@ -1,36 +1,58 @@ package ldapi import ( + "context" "encoding/json" + "errors" "fmt" "net/http" "net/url" "strconv" + retryablehttp "github.com/hashicorp/go-retryablehttp" ldapi "github.com/launchdarkly/api-client-go/v15" lcr "github.com/launchdarkly/find-code-references-in-pull-request/config" gha "github.com/launchdarkly/find-code-references-in-pull-request/internal/github_actions" "github.com/launchdarkly/find-code-references-in-pull-request/internal/version" - "github.com/pkg/errors" ) -func GetAllFlags(config *lcr.Config) ([]ldapi.FeatureFlag, error) { +const pageSize = 100 + +type flagPage struct { + items []ldapi.FeatureFlag + headers http.Header +} + +func GetAllFlags(ctx context.Context, config *lcr.Config) (flags []ldapi.FeatureFlag, err error) { + inventoryContext, cancel := context.WithTimeout(ctx, inventoryTimeout) + defer cancel() + + fetcher := newFlagFetcher(inventoryContext, fetcherOptions{ + log: func(format string, args ...any) { + gha.Log(format, args...) + }, + }) + fetcher.diagnostics.emitStart() + defer func() { + fetcher.diagnostics.emitSummary(fetcher.clock(), fetcher.outcome(err), err == nil) + }() + gha.Debug("Fetching all flags for project") params := url.Values{} params.Add("env", config.LdEnvironment) - activeFlags, err := getFlags(config, params) + activeFlags, err := fetcher.getFlags(config, params, flagCollectionActive, config.IncludeArchivedFlags) if err != nil { return []ldapi.FeatureFlag{}, err } - flags := make([]ldapi.FeatureFlag, 0, len(activeFlags)) + flags = make([]ldapi.FeatureFlag, 0, len(activeFlags)) flags = append(flags, activeFlags...) if config.IncludeArchivedFlags { params.Add("filter", "state:archived") - archivedFlags, err := getFlags(config, params) - if err != nil { - return []ldapi.FeatureFlag{}, err + archivedFlags, archivedErr := fetcher.getFlags(config, params, flagCollectionArchived, false) + if archivedErr != nil { + return []ldapi.FeatureFlag{}, archivedErr } flags = append(flags, archivedFlags...) } @@ -39,62 +61,168 @@ func GetAllFlags(config *lcr.Config) ([]ldapi.FeatureFlag, error) { return flags, nil } -func getFlags(config *lcr.Config, params url.Values) ([]ldapi.FeatureFlag, error) { +func (f *flagFetcher) getFlags( + config *lcr.Config, + params url.Values, + collection string, + nextCollection bool, +) ([]ldapi.FeatureFlag, error) { pageParams := make(url.Values, len(params)+2) for key, values := range params { pageParams[key] = append([]string(nil), values...) } - const pageSize = 100 pageParams.Set("limit", strconv.Itoa(pageSize)) - client := &http.Client{} flags := []ldapi.FeatureFlag{} for offset := 0; ; offset += pageSize { + f.collection = collection + f.offset = offset + f.attempt = 0 + f.waitingRetry = false + f.pendingRetryWait = 0 + if err := f.waitForNextPage(); err != nil { + f.diagnostics.recordFailure(collection, offset) + return []ldapi.FeatureFlag{}, err + } pageParams.Set("offset", strconv.Itoa(offset)) - page, err := getFlagPage(client, config, pageParams) + page, err := f.getFlagPage(config, pageParams) if err != nil { + f.diagnostics.recordFailure(collection, offset) return []ldapi.FeatureFlag{}, err } - // The migration guide uses an empty page to signal completion. - if len(page) == 0 { + + f.diagnostics.recordPage(collection, len(page.items)) + if len(page.items) == 0 { + if nextCollection { + f.scheduleProactiveWait(page.headers, collection, offset) + } return flags, nil } - flags = append(flags, page...) + + flags = append(flags, page.items...) + f.scheduleProactiveWait(page.headers, collection, offset) } } -func getFlagPage(client *http.Client, config *lcr.Config, params url.Values) ([]ldapi.FeatureFlag, error) { - url := fmt.Sprintf("%s/api/v2/flags/%s", config.LdInstance, config.LdProject) - req, err := http.NewRequest(http.MethodGet, url, nil) - if err != nil { - return []ldapi.FeatureFlag{}, err +func (f *flagFetcher) getFlagPage(config *lcr.Config, params url.Values) (flagPage, error) { + if err := f.contextError(); err != nil { + return flagPage{}, err } - req.URL.RawQuery = params.Encode() - req.Header.Add("Authorization", config.ApiToken) - req.Header.Add("LD-API-Version", "20240415") - req.Header.Add("User-Agent", fmt.Sprintf("find-code-references-pr/%s", version.Version)) - resp, err := client.Do(req) + endpoint := fmt.Sprintf("%s/api/v2/flags/%s", config.LdInstance, config.LdProject) + retryRequest, err := retryablehttp.NewRequestWithContext(f.ctx, http.MethodGet, endpoint, nil) if err != nil { - return []ldapi.FeatureFlag{}, err + return flagPage{}, f.newError(reasonRequest, 0, nil) } - defer resp.Body.Close() - decoder := json.NewDecoder(resp.Body) + retryRequest.URL.RawQuery = params.Encode() + retryRequest.Header.Set("Authorization", config.ApiToken) + retryRequest.Header.Set("LD-API-Version", "20240415") + retryRequest.Header.Set("User-Agent", fmt.Sprintf("find-code-references-pr/%s", version.Version)) - if resp.StatusCode != http.StatusOK { - var r interface{} - if err := decoder.Decode(&r); err != nil { - return []ldapi.FeatureFlag{}, errors.Wrapf(err, "unexpected status code: %d. unable to parse response", resp.StatusCode) + response, requestErr := f.client.Do(retryRequest) + if requestErr != nil { + closeResponseBody(response) + if existing := new(flagFetchError); errors.As(requestErr, &existing) { + return flagPage{}, requestErr } - err := fmt.Errorf("unexpected status code: %d with response: %#v", resp.StatusCode, r) - return []ldapi.FeatureFlag{}, err + return flagPage{}, f.wrapRequestError(requestErr) + } + if response == nil { + return flagPage{}, f.newError(reasonTransport, 0, nil) } - flags := ldapi.FeatureFlags{} - err = decoder.Decode(&flags) - if err != nil { - return []ldapi.FeatureFlag{}, err + if response.StatusCode != http.StatusOK { + status := response.StatusCode + closeResponseBody(response) + if status == http.StatusTooManyRequests { + return flagPage{}, f.newError(reasonRateLimitExhausted, status, nil) + } + return flagPage{}, f.newError(reasonHTTP, status, nil) + } + if response.Body == nil { + closeResponseBody(response) + return flagPage{}, f.newError(reasonDecode, response.StatusCode, nil) + } + + var flags ldapi.FeatureFlags + decodeErr := json.NewDecoder(response.Body).Decode(&flags) + headers := response.Header.Clone() + closeResponseBody(response) + if decodeErr != nil { + return flagPage{}, f.wrapDecodeError(decodeErr) + } + return flagPage{items: flags.Items, headers: headers}, nil +} + +func (f *flagFetcher) waitForNextPage() error { + if !f.hasNotBefore { + return nil + } + if f.notBeforeUnfulfillable { + f.hasNotBefore = false + f.notBeforeUnfulfillable = false + return f.newError(reasonRateLimitDeadline, 0, context.DeadlineExceeded) } - return flags.Items, nil + now := f.clock() + wait := f.notBefore.Sub(now) + if wait <= 0 { + f.hasNotBefore = false + f.notBeforeUnfulfillable = false + return nil + } + if deadline, ok := f.ctx.Deadline(); !ok || fitsDeadline(now, &deadline, wait) { + f.waitingProactive = true + waitErr := f.wait(f.ctx, wait) + f.waitingProactive = false + f.hasNotBefore = false + f.notBeforeUnfulfillable = false + if waitErr != nil { + if contextErr := f.ctx.Err(); contextErr != nil { + reason := reasonDeadline + if errors.Is(contextErr, context.Canceled) { + reason = reasonCanceled + } + if errors.Is(contextErr, context.DeadlineExceeded) { + reason = reasonRateLimitDeadline + } + return f.newError(reason, 0, contextErr) + } + return f.newError(reasonTransport, 0, waitErr) + } + return nil + } + + f.hasNotBefore = false + f.notBeforeUnfulfillable = false + return f.newError(reasonRateLimitDeadline, 0, context.DeadlineExceeded) +} + +func (f *flagFetcher) scheduleProactiveWait(headers http.Header, collection string, offset int) { + deadline, hasDeadline := f.ctx.Deadline() + decision := proactiveWait( + f.clock(), + deadlinePointer(deadline, hasDeadline), + parseRateLimitHeaders(headers), + f.jitter, + ) + if decision.unfulfillable { + f.notBeforeUnfulfillable = true + f.hasNotBefore = true + return + } + if decision.delay <= 0 { + return + } + + f.notBefore = f.clock().Add(decision.delay) + f.hasNotBefore = true + f.diagnostics.recordProactiveWait(decision.delay) + f.diagnostics.emitWait(collection, offset, 1, decision.reason, decision.delay) +} + +func closeResponseBody(response *http.Response) { + if response != nil && response.Body != nil { + _ = response.Body.Close() + } } diff --git a/internal/ldclient/flags_test.go b/internal/ldclient/flags_test.go index 841e78b3..00c0fcec 100644 --- a/internal/ldclient/flags_test.go +++ b/internal/ldclient/flags_test.go @@ -1,6 +1,7 @@ package ldapi import ( + "context" "encoding/json" "fmt" "net/http" @@ -60,7 +61,7 @@ func TestGetAllFlagsPagination(t *testing.T) { } config := serveFlagPages(t, tt.archivedPages != nil, pages) - flags, err := GetAllFlags(config) + flags, err := GetAllFlags(context.Background(), config) require.NoError(t, err) gotKeys := make([]string, len(flags)) @@ -80,7 +81,7 @@ func TestGetAllFlagsPreservesMetadata(t *testing.T) { {offset: 100, body: `{"items":[]}`}, }) - flags, err := GetAllFlags(config) + flags, err := GetAllFlags(context.Background(), config) require.NoError(t, err) require.Len(t, flags, 1) @@ -120,7 +121,8 @@ func TestGetFlagsPreservesQueryParameters(t *testing.T) { "offset": {"42"}, } - flags, err := getFlags(config, params) + fetcher, _ := newTestFetcher(t) + flags, err := fetcher.getFlags(config, params, flagCollectionArchived, false) require.NoError(t, err) assert.Len(t, flags, 1) @@ -144,11 +146,11 @@ func TestGetAllFlagsErrors(t *testing.T) { }, { name: "non-JSON error", status: http.StatusBadGateway, - body: "bad gateway", wantError: "502. unable to parse response", + body: "bad gateway", wantError: "status=502 reason=http_error", }, { name: "invalid success JSON", status: http.StatusOK, - body: "not JSON", wantError: "invalid character", + body: "not JSON", wantError: "status=200 reason=decode_error", }, } for _, tt := range tests { @@ -173,7 +175,7 @@ func TestGetAllFlagsErrors(t *testing.T) { pages = append(pages, flagPageResponse{offset: offset, filter: filter, status: tt.status, body: tt.body}) config := serveFlagPages(t, archived, pages) - flags, err := GetAllFlags(config) + flags, err := GetAllFlags(context.Background(), config) require.ErrorContains(t, err, tt.wantError) assert.Empty(t, flags) @@ -184,10 +186,11 @@ func TestGetAllFlagsErrors(t *testing.T) { } type flagPageResponse struct { - offset int - filter string - status int - body string + offset int + filter string + status int + body string + headers http.Header } func serveFlagPages(t *testing.T, includeArchived bool, pages []flagPageResponse) *lcr.Config { @@ -215,6 +218,11 @@ func serveFlagPages(t *testing.T, includeArchived bool, pages []flagPageResponse } assert.Equal(t, query, r.URL.Query()) w.Header().Set("Content-Type", "application/json") + for name, values := range page.headers { + for _, value := range values { + w.Header().Add(name, value) + } + } status := page.status if status == 0 { status = http.StatusOK diff --git a/internal/ldclient/retry.go b/internal/ldclient/retry.go new file mode 100644 index 00000000..5ef5e050 --- /dev/null +++ b/internal/ldclient/retry.go @@ -0,0 +1,727 @@ +package ldapi + +import ( + "context" + "errors" + "math/rand" + "net" + "net/http" + "strconv" + "strings" + "time" + + retryablehttp "github.com/hashicorp/go-retryablehttp" +) + +const ( + inventoryTimeout = 120 * time.Second + attemptTimeout = 30 * time.Second + maxRetries = 4 + + flagCollectionActive = "active" + flagCollectionArchived = "archived" +) + +const maxInt64Value = int64(1<<63 - 1) + +type fetcherOptions struct { + clock func() time.Time + jitter func(time.Duration, time.Duration) time.Duration + wait func(context.Context, time.Duration) error + httpClient *http.Client + log diagnosticLog +} + +type flagFetcher struct { + ctx context.Context + + client *retryablehttp.Client + clock func() time.Time + jitter func(time.Duration, time.Duration) time.Duration + wait func(context.Context, time.Duration) error + diagnostics *fetchDiagnostics + + collection string + offset int + attempt int + + pendingRetryWait time.Duration + waitingRetry bool + waitingProactive bool + notBefore time.Time + hasNotBefore bool + notBeforeUnfulfillable bool +} + +func newFlagFetcher(ctx context.Context, options fetcherOptions) *flagFetcher { + if options.clock == nil { + options.clock = time.Now + } + if options.jitter == nil { + options.jitter = uniformJitter + } + if options.wait == nil { + options.wait = waitWithContext + } + + client := retryablehttp.NewClient() + if options.httpClient != nil { + client.HTTPClient = options.httpClient + } + client.HTTPClient.Timeout = attemptTimeout + client.Logger = nil + client.RetryMax = maxRetries + fetcher := &flagFetcher{ + ctx: ctx, + client: client, + clock: options.clock, + jitter: options.jitter, + wait: options.wait, + diagnostics: newFetchDiagnostics(options.clock(), options.log), + } + client.CheckRetry = fetcher.checkRetry + client.Backoff = fetcher.retryBackoff + client.RequestLogHook = fetcher.requestLogHook + client.ResponseLogHook = fetcher.responseLogHook + client.ErrorHandler = retryablehttp.PassthroughErrorHandler + return fetcher +} + +func (f *flagFetcher) requestLogHook(_ retryablehttp.Logger, _ *http.Request, retryNumber int) { + f.attempt = retryNumber + 1 + f.waitingRetry = false + f.pendingRetryWait = 0 + f.diagnostics.recordAttempt(retryNumber) +} + +func (f *flagFetcher) responseLogHook(_ retryablehttp.Logger, response *http.Response) { + if response != nil { + f.diagnostics.recordStatus(response.StatusCode) + } +} + +func (f *flagFetcher) retryBackoff(_ time.Duration, _ time.Duration, _ int, _ *http.Response) time.Duration { + wait := f.pendingRetryWait + f.pendingRetryWait = 0 + return wait +} + +func (f *flagFetcher) checkRetry(ctx context.Context, response *http.Response, requestErr error) (bool, error) { + if response != nil { + f.diagnostics.recordStatus(response.StatusCode) + } + + isRateLimited := response != nil && response.StatusCode == http.StatusTooManyRequests + if isRateLimited { + hints := parseRateLimitHeaders(response.Header) + f.diagnostics.responses429++ + f.diagnostics.emitRateLimit(f.collection, f.offset, f.currentAttempt(), response.StatusCode, hints) + } + + if ctx.Err() != nil { + return false, ctx.Err() + } + if requestErr != nil { + return false, nil + } + if !isRateLimited { + return false, nil + } + + attempt := f.currentAttempt() + if attempt > maxRetries+1 { + return true, nil + } + if attempt == maxRetries+1 { + // Keep shouldRetry true so PassthroughErrorHandler returns the final 429. + return true, nil + } + + hints := parseRateLimitHeaders(response.Header) + deadline, hasDeadline := f.ctx.Deadline() + if !hasDeadline { + deadline = time.Time{} + } + decision := retryWait( + f.clock(), + deadlinePointer(deadline, hasDeadline), + attempt-1, + hints, + f.jitter, + ) + if decision.unfulfillable { + return false, f.newError(reasonRateLimitDeadline, http.StatusTooManyRequests, context.DeadlineExceeded) + } + + f.pendingRetryWait = decision.delay + f.waitingRetry = true + f.diagnostics.recordRetryWait(decision.delay) + f.diagnostics.emitWait(f.collection, f.offset, attempt+1, decision.reason, decision.delay) + return true, nil +} + +func (f *flagFetcher) currentAttempt() int { + if f.attempt > 0 { + return f.attempt + } + return 1 +} + +type fetchErrorReason string + +const ( + reasonCanceled fetchErrorReason = "canceled" + reasonDeadline fetchErrorReason = "deadline" + reasonDecode fetchErrorReason = "decode_error" + reasonHTTP fetchErrorReason = "http_error" + reasonRateLimitDeadline fetchErrorReason = "rate_limit_deadline" + reasonRateLimitExhausted fetchErrorReason = "rate_limit_exhausted" + reasonRequest fetchErrorReason = "request_error" + reasonTransport fetchErrorReason = "transport_error" +) + +type flagFetchError struct { + reason fetchErrorReason + collection string + offset int + attempt int + status int + cause error +} + +func (e *flagFetchError) Error() string { + message := "feature flag fetch failed" + if e.collection != "" { + message += " collection=" + e.collection + } + message += " offset=" + strconv.Itoa(e.offset) + if e.attempt > 0 { + message += " attempt=" + strconv.Itoa(e.attempt) + } + if e.status > 0 { + message += " status=" + strconv.Itoa(e.status) + } + message += " reason=" + string(e.reason) + return message +} + +func (e *flagFetchError) Unwrap() error { + return e.cause +} + +func (f *flagFetcher) newError(reason fetchErrorReason, status int, cause error) *flagFetchError { + if reason == reasonRateLimitDeadline && cause == nil { + cause = context.DeadlineExceeded + } + return &flagFetchError{ + reason: reason, + collection: f.collection, + offset: f.offset, + attempt: f.currentAttempt(), + status: status, + cause: cause, + } +} + +func (f *flagFetcher) contextError() error { + contextErr := f.ctx.Err() + if contextErr == nil { + return nil + } + reason := reasonDeadline + if errors.Is(contextErr, context.Canceled) { + reason = reasonCanceled + } + if (f.waitingRetry || f.waitingProactive) && errors.Is(contextErr, context.DeadlineExceeded) { + reason = reasonRateLimitDeadline + } + return f.newError(reason, 0, contextErr) +} + +func (f *flagFetcher) wrapRequestError(requestErr error) error { + if requestErr == nil { + return nil + } + if existing := new(flagFetchError); errors.As(requestErr, &existing) { + return requestErr + } + if contextErr := f.ctx.Err(); contextErr != nil { + reason := reasonDeadline + if errors.Is(contextErr, context.Canceled) { + reason = reasonCanceled + } + if (f.waitingRetry || f.waitingProactive) && errors.Is(contextErr, context.DeadlineExceeded) { + reason = reasonRateLimitDeadline + } + return f.newError(reason, 0, requestErr) + } + return f.newError(reasonTransport, 0, requestErr) +} + +func (f *flagFetcher) wrapDecodeError(decodeErr error) *flagFetchError { + reason := reasonDecode + cause := decodeErr + + switch { + case errors.Is(decodeErr, context.Canceled): + reason = reasonCanceled + case errors.Is(decodeErr, context.DeadlineExceeded): + reason = reasonDeadline + default: + var timeoutErr net.Error + if errors.As(decodeErr, &timeoutErr) && timeoutErr.Timeout() { + reason = reasonDeadline + } else if contextErr := f.ctx.Err(); contextErr != nil { + cause = errors.Join(decodeErr, contextErr) + if errors.Is(contextErr, context.Canceled) { + reason = reasonCanceled + } else { + reason = reasonDeadline + } + } + } + + return f.newError(reason, http.StatusOK, cause) +} + +func (f *flagFetcher) outcome(err error) string { + if err == nil { + return "success" + } + var fetchErr *flagFetchError + if errors.As(err, &fetchErr) { + switch fetchErr.reason { + case reasonCanceled: + return "canceled" + case reasonDeadline: + return "deadline" + case reasonDecode: + return "decode_error" + case reasonHTTP: + return "http_error" + case reasonRateLimitDeadline: + return "rate_limit_deadline" + case reasonRateLimitExhausted: + return "rate_limit_exhausted" + case reasonTransport, reasonRequest: + return "transport_error" + } + } + if errors.Is(err, context.Canceled) { + return "canceled" + } + if errors.Is(err, context.DeadlineExceeded) { + return "deadline" + } + return "transport_error" +} + +func deadlinePointer(deadline time.Time, ok bool) *time.Time { + if !ok { + return nil + } + return &deadline +} + +func waitWithContext(ctx context.Context, wait time.Duration) error { + if wait <= 0 { + return nil + } + timer := time.NewTimer(wait) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + +func uniformJitter(minimum, maximum time.Duration) time.Duration { + if maximum <= minimum { + return minimum + } + return minimum + time.Duration(rand.Int63n(int64(maximum-minimum)+1)) +} + +type parsedRateLimitHeaders struct { + globalResetValues []int64 + authTokenResetValues []int64 + globalResetOverflow bool + authResetOverflow bool + + retryAfterSeconds []uint64 + retryAfterDates []time.Time + retryAfterOverflow bool + + globalRemaining *int64 + routeRemaining *int64 + authTokenRemaining *int64 +} + +func parseRateLimitHeaders(headers http.Header) parsedRateLimitHeaders { + parsed := parsedRateLimitHeaders{} + for _, value := range headers.Values("X-Ratelimit-Reset") { + parsedValue, valid, overflow := parseUnsignedHeader(value) + if overflow || valid && parsedValue > uint64(maxInt64Value) { + parsed.globalResetOverflow = true + } + if valid && parsedValue <= uint64(maxInt64Value) { + parsed.globalResetValues = append(parsed.globalResetValues, int64(parsedValue)) + } + } + for _, value := range headers.Values("X-Ratelimit-Auth-Token-Reset") { + parsedValue, valid, overflow := parseUnsignedHeader(value) + if overflow || valid && parsedValue > uint64(maxInt64Value) { + parsed.authResetOverflow = true + } + if valid && parsedValue <= uint64(maxInt64Value) { + parsed.authTokenResetValues = append(parsed.authTokenResetValues, int64(parsedValue)) + } + } + for _, value := range headers.Values("Retry-After") { + trimmed := strings.TrimSpace(value) + parsedValue, valid, overflow := parseUnsignedHeader(trimmed) + if overflow { + parsed.retryAfterOverflow = true + } + if valid { + parsed.retryAfterSeconds = append(parsed.retryAfterSeconds, parsedValue) + continue + } + if date, err := http.ParseTime(trimmed); err == nil { + parsed.retryAfterDates = append(parsed.retryAfterDates, date) + } + } + parsed.globalRemaining = parseRemainingHeader(headers, "X-Ratelimit-Global-Remaining") + parsed.routeRemaining = parseRemainingHeader(headers, "X-Ratelimit-Route-Remaining") + parsed.authTokenRemaining = parseRemainingHeader(headers, "X-Ratelimit-Auth-Token-Remaining") + return parsed +} + +func parseUnsignedHeader(value string) (uint64, bool, bool) { + trimmed := strings.TrimSpace(value) + if trimmed == "" || strings.HasPrefix(trimmed, "-") { + return 0, false, false + } + parsed, err := strconv.ParseUint(trimmed, 10, 64) + if err == nil { + return parsed, true, false + } + if errors.Is(err, strconv.ErrRange) { + return 0, false, true + } + return 0, false, false +} + +func parseRemainingHeader(headers http.Header, name string) *int64 { + for _, value := range headers.Values(name) { + parsed, valid, _ := parseUnsignedHeader(value) + if valid && parsed <= uint64(maxInt64Value) { + converted := int64(parsed) + return &converted + } + } + return nil +} + +func (h parsedRateLimitHeaders) globalResetMillis() *int64 { + return maxInt64Pointer(h.globalResetValues) +} + +func (h parsedRateLimitHeaders) authTokenResetMillis() *int64 { + return maxInt64Pointer(h.authTokenResetValues) +} + +func (h parsedRateLimitHeaders) retryAfterSecondsValue() *int64 { + var selected *int64 + for _, value := range h.retryAfterSeconds { + if value > uint64(maxInt64Value) { + continue + } + converted := int64(value) + if selected == nil || converted > *selected { + selected = &converted + } + } + return selected +} + +func (h parsedRateLimitHeaders) retryAfterDateMillis() *int64 { + var selected *int64 + for _, date := range h.retryAfterDates { + converted := date.UnixMilli() + if selected == nil || converted > *selected { + selected = &converted + } + } + return selected +} + +func (h parsedRateLimitHeaders) globalRemainingValue() *int64 { + return h.globalRemaining +} + +func (h parsedRateLimitHeaders) routeRemainingValue() *int64 { + return h.routeRemaining +} + +func (h parsedRateLimitHeaders) authTokenRemainingValue() *int64 { + return h.authTokenRemaining +} + +func maxInt64Pointer(values []int64) *int64 { + if len(values) == 0 { + return nil + } + maximum := values[0] + for _, value := range values[1:] { + if value > maximum { + maximum = value + } + } + return &maximum +} + +type waitDecision struct { + delay time.Duration + reason string + unfulfillable bool +} + +func retryWait( + now time.Time, + deadline *time.Time, + retryNumber int, + hints parsedRateLimitHeaders, + jitter func(time.Duration, time.Duration) time.Duration, +) waitDecision { + latest, hasFuture, unfulfillable := latestFutureHint(now, deadline, hints) + if unfulfillable { + return waitDecision{reason: "rate_limit_reset", unfulfillable: true} + } + + if hasFuture { + wait := latest.Sub(now) + jitterWait := boundedJitter(jitter, 100*time.Millisecond, time.Second) + if wait > time.Duration(maxInt64Value)-jitterWait { + return waitDecision{reason: "rate_limit_reset", unfulfillable: true} + } + wait += jitterWait + if !fitsDeadline(now, deadline, wait) { + return waitDecision{reason: "rate_limit_reset", unfulfillable: true} + } + return waitDecision{delay: wait, reason: "rate_limit_reset"} + } + + window := retryWindow(retryNumber) + wait := boundedJitter(jitter, window/2, window) + if !fitsDeadline(now, deadline, wait) { + return waitDecision{reason: "rate_limit_backoff", unfulfillable: true} + } + return waitDecision{delay: wait, reason: "rate_limit_backoff"} +} + +func latestFutureHint(now time.Time, deadline *time.Time, hints parsedRateLimitHeaders) (time.Time, bool, bool) { + latest, hasFuture, unfulfillable := latestRetryAfterHint(now, deadline, hints) + if unfulfillable { + return time.Time{}, false, true + } + consider := func(at time.Time) bool { + if !at.After(now) { + return true + } + if deadline != nil && at.After(*deadline) { + return false + } + if !hasFuture || at.After(latest) { + latest = at + hasFuture = true + } + return true + } + considerResets := func(values []int64, overflow bool) bool { + if overflow { + return false + } + for _, millis := range values { + if deadline != nil && millis > deadline.UnixMilli() { + return false + } + if !consider(time.UnixMilli(millis)) { + return false + } + } + return true + } + + sharedResetApplicable := (hints.globalRemaining != nil && *hints.globalRemaining == 0) || + (hints.routeRemaining != nil && *hints.routeRemaining == 0) || + (hints.globalRemaining == nil && hints.routeRemaining == nil) + if sharedResetApplicable && !considerResets(hints.globalResetValues, hints.globalResetOverflow) { + return time.Time{}, false, true + } + tokenResetApplicable := hints.authTokenRemaining == nil || *hints.authTokenRemaining == 0 + if tokenResetApplicable && !considerResets(hints.authTokenResetValues, hints.authResetOverflow) { + return time.Time{}, false, true + } + return latest, hasFuture, false +} + +func retryWindow(retryNumber int) time.Duration { + window := 2 * time.Second + for i := 0; i < retryNumber && window < 30*time.Second; i++ { + if window > 15*time.Second { + return 30 * time.Second + } + window *= 2 + } + if window > 30*time.Second { + return 30 * time.Second + } + return window +} + +func fitsDeadline(now time.Time, deadline *time.Time, wait time.Duration) bool { + if deadline == nil { + return true + } + remaining := deadline.Sub(now) + return wait >= 0 && wait < remaining +} + +func boundedJitter(jitter func(time.Duration, time.Duration) time.Duration, minimum, maximum time.Duration) time.Duration { + value := jitter(minimum, maximum) + if value < minimum { + return minimum + } + if value > maximum { + return maximum + } + return value +} + +func proactiveWait( + now time.Time, + deadline *time.Time, + hints parsedRateLimitHeaders, + jitter func(time.Duration, time.Duration) time.Duration, +) waitDecision { + var wait time.Duration + serviceDelay := false + + considerQuota := func(remaining *int64, resetValues []int64, resetOverflow bool) bool { + if remaining == nil || *remaining != 0 { + return true + } + if resetOverflow { + return false + } + if len(resetValues) == 0 { + fallback := boundedJitter(jitter, time.Second, 2*time.Second) + if fallback > wait { + wait = fallback + } + return true + } + var futureReset time.Time + for _, millis := range resetValues { + if deadline != nil && millis > deadline.UnixMilli() { + return false + } + reset := time.UnixMilli(millis) + if reset.After(now) && (futureReset.IsZero() || reset.After(futureReset)) { + futureReset = reset + } + } + if futureReset.IsZero() { + return true + } + resetWait := futureReset.Sub(now) + if resetWait > wait { + wait = resetWait + } + serviceDelay = true + return true + } + + if !considerQuota(hints.globalRemaining, hints.globalResetValues, hints.globalResetOverflow) { + return waitDecision{reason: "proactive_rate_limit", unfulfillable: true} + } + if !considerQuota(hints.routeRemaining, hints.globalResetValues, hints.globalResetOverflow) { + return waitDecision{reason: "proactive_rate_limit", unfulfillable: true} + } + if !considerQuota(hints.authTokenRemaining, hints.authTokenResetValues, hints.authResetOverflow) { + return waitDecision{reason: "proactive_rate_limit", unfulfillable: true} + } + + if hints.globalRemaining != nil && *hints.globalRemaining == 0 || + hints.routeRemaining != nil && *hints.routeRemaining == 0 || + hints.authTokenRemaining != nil && *hints.authTokenRemaining == 0 { + latest, hasFuture, unfulfillable := latestRetryAfterHint(now, deadline, hints) + if unfulfillable { + return waitDecision{reason: "proactive_rate_limit", unfulfillable: true} + } + if hasFuture { + retryAfterWait := latest.Sub(now) + if retryAfterWait > wait { + wait = retryAfterWait + } + serviceDelay = serviceDelay || retryAfterWait > 0 + } + } + + if wait <= 0 { + return waitDecision{reason: "proactive_rate_limit"} + } + if serviceDelay { + jitterWait := boundedJitter(jitter, 100*time.Millisecond, time.Second) + if wait > time.Duration(maxInt64Value)-jitterWait { + return waitDecision{reason: "proactive_rate_limit", unfulfillable: true} + } + wait += jitterWait + } + if !fitsDeadline(now, deadline, wait) { + return waitDecision{reason: "proactive_rate_limit", unfulfillable: true} + } + return waitDecision{delay: wait, reason: "proactive_rate_limit"} +} + +func latestRetryAfterHint(now time.Time, deadline *time.Time, hints parsedRateLimitHeaders) (time.Time, bool, bool) { + if hints.retryAfterOverflow { + return time.Time{}, false, true + } + var latest time.Time + hasFuture := false + consider := func(at time.Time) bool { + if !at.After(now) { + return true + } + if deadline != nil && at.After(*deadline) { + return false + } + if !hasFuture || at.After(latest) { + latest = at + hasFuture = true + } + return true + } + for _, seconds := range hints.retryAfterSeconds { + if seconds > uint64(maxInt64Value)/uint64(time.Second) { + return time.Time{}, false, true + } + wait := time.Duration(seconds) * time.Second + if deadline != nil && !fitsDeadline(now, deadline, wait) { + return time.Time{}, false, true + } + if !consider(now.Add(wait)) { + return time.Time{}, false, true + } + } + for _, date := range hints.retryAfterDates { + if !consider(date) { + return time.Time{}, false, true + } + } + return latest, hasFuture, false +} diff --git a/internal/ldclient/retry_test.go b/internal/ldclient/retry_test.go new file mode 100644 index 00000000..0824db01 --- /dev/null +++ b/internal/ldclient/retry_test.go @@ -0,0 +1,1796 @@ +package ldapi + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/url" + "os" + "strconv" + "strings" + "sync/atomic" + "testing" + "testing/iotest" + "time" + + ldapi "github.com/launchdarkly/api-client-go/v15" + lcr "github.com/launchdarkly/find-code-references-in-pull-request/config" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestRetryWaitUsesLatestServiceHintAndJitter(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + hints := parseRateLimitHeaders(http.Header{ + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(3*time.Second).UnixMilli(), 10)}, + "X-Ratelimit-Auth-Token-Reset": {strconv.FormatInt(now.Add(5*time.Second).UnixMilli(), 10)}, + "Retry-After": {"2"}, + }) + + decision := retryWait(now, &deadline, 0, hints, func(_, maximum time.Duration) time.Duration { + return maximum + }) + + assert.False(t, decision.unfulfillable) + assert.Equal(t, "rate_limit_reset", decision.reason) + assert.Equal(t, 6*time.Second, decision.delay) +} + +func TestRetryWaitUsesApplicableQuotaResets(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + maximumJitter := func(_, maximum time.Duration) time.Duration { return maximum } + tests := []struct { + name string + global string + route string + token string + sharedReset string + tokenReset string + retryAfter string + deadlineOffset time.Duration + wantDelay time.Duration + wantReason string + wantUnfulfill bool + }{ + { + name: "route exhausted, token available", global: "1", route: "0", token: "5", + sharedReset: "2", tokenReset: "8", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "global exhausted, token available", global: "0", route: "1", token: "5", + sharedReset: "2", tokenReset: "8", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "token exhausted, shared quotas available", global: "1", route: "1", token: "0", + sharedReset: "8", tokenReset: "2", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "only route remaining supplied", global: "?", route: "1", token: "0", + sharedReset: "8", tokenReset: "2", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "only global remaining supplied", global: "1", route: "?", token: "0", + sharedReset: "8", tokenReset: "2", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "both reset families exhausted", global: "1", route: "0", token: "0", + sharedReset: "2", tokenReset: "4", deadlineOffset: 10 * time.Second, + wantDelay: 5 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "unknown token count", global: "1", route: "0", token: "?", + sharedReset: "2", tokenReset: "8", deadlineOffset: 10 * time.Second, + wantDelay: 9 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "only reset headers supplied", global: "?", route: "?", token: "?", + sharedReset: "2", tokenReset: "4", deadlineOffset: 10 * time.Second, + wantDelay: 5 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "all known quotas have capacity", global: "1", route: "1", token: "5", + sharedReset: "8", tokenReset: "8", deadlineOffset: 5 * time.Second, + wantDelay: 2 * time.Second, wantReason: "rate_limit_backoff", + }, + { + name: "retry-after exceeds relevant reset", global: "1", route: "0", token: "5", + sharedReset: "2", tokenReset: "8", retryAfter: "3", deadlineOffset: 5 * time.Second, + wantDelay: 4 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "jitter reaches deadline", global: "1", route: "0", token: "5", + sharedReset: "2", tokenReset: "8", retryAfter: "4", deadlineOffset: 5 * time.Second, + wantReason: "rate_limit_reset", wantUnfulfill: true, + }, + { + name: "retry-after exceeds deadline", global: "1", route: "0", token: "5", + sharedReset: "2", tokenReset: "8", retryAfter: "6", deadlineOffset: 5 * time.Second, + wantReason: "rate_limit_reset", wantUnfulfill: true, + }, + { + name: "ignored token reset overflows", global: "1", route: "0", token: "5", + sharedReset: "2", tokenReset: "9223372036854775808", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "ignored shared reset overflows", global: "1", route: "1", token: "0", + sharedReset: "9223372036854775808", tokenReset: "2", deadlineOffset: 5 * time.Second, + wantDelay: 3 * time.Second, wantReason: "rate_limit_reset", + }, + { + name: "exhausted token reset overflows", global: "1", route: "0", token: "0", + sharedReset: "2", tokenReset: "9223372036854775808", deadlineOffset: 5 * time.Second, + wantReason: "rate_limit_reset", wantUnfulfill: true, + }, + { + name: "unknown token reset overflows", global: "1", route: "0", token: "?", + sharedReset: "2", tokenReset: "9223372036854775808", deadlineOffset: 5 * time.Second, + wantReason: "rate_limit_reset", wantUnfulfill: true, + }, + { + name: "exhausted route reset overflows", global: "1", route: "0", token: "5", + sharedReset: "9223372036854775808", tokenReset: "8", deadlineOffset: 5 * time.Second, + wantReason: "rate_limit_reset", wantUnfulfill: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + headers := http.Header{} + if tt.global != "?" { + headers.Set("X-Ratelimit-Global-Remaining", tt.global) + } + if tt.route != "?" { + headers.Set("X-Ratelimit-Route-Remaining", tt.route) + } + if tt.token != "?" { + headers.Set("X-Ratelimit-Auth-Token-Remaining", tt.token) + } + if tt.sharedReset != "" { + if strings.HasPrefix(tt.sharedReset, "9223372036854775808") { + headers.Set("X-Ratelimit-Reset", tt.sharedReset) + } else { + offset, err := time.ParseDuration(tt.sharedReset + "s") + require.NoError(t, err) + headers.Set("X-Ratelimit-Reset", strconv.FormatInt(now.Add(offset).UnixMilli(), 10)) + } + } + if tt.tokenReset != "" { + if strings.HasPrefix(tt.tokenReset, "9223372036854775808") { + headers.Set("X-Ratelimit-Auth-Token-Reset", tt.tokenReset) + } else { + offset, err := time.ParseDuration(tt.tokenReset + "s") + require.NoError(t, err) + headers.Set("X-Ratelimit-Auth-Token-Reset", strconv.FormatInt(now.Add(offset).UnixMilli(), 10)) + } + } + if tt.retryAfter != "" { + headers.Set("Retry-After", tt.retryAfter) + } + + decision := retryWait( + now, + timePtr(now.Add(tt.deadlineOffset)), + 0, + parseRateLimitHeaders(headers), + maximumJitter, + ) + assert.Equal(t, tt.wantReason, decision.reason) + assert.Equal(t, tt.wantUnfulfill, decision.unfulfillable) + if tt.wantUnfulfill { + assert.Zero(t, decision.delay) + } else { + assert.Equal(t, tt.wantDelay, decision.delay) + } + }) + } + + for _, remaining := range []string{"?", "malformed"} { + t.Run("unknown token "+remaining, func(t *testing.T) { + headers := http.Header{ + "X-Ratelimit-Global-Remaining": {"1"}, + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Auth-Token-Reset": { + strconv.FormatInt(now.Add(8*time.Second).UnixMilli(), 10), + }, + "X-Ratelimit-Reset": { + strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10), + }, + } + if remaining == "malformed" { + headers.Set("X-Ratelimit-Auth-Token-Remaining", remaining) + } + decision := retryWait( + now, + timePtr(now.Add(10*time.Second)), + 0, + parseRateLimitHeaders(headers), + maximumJitter, + ) + assert.Equal(t, 9*time.Second, decision.delay) + assert.Equal(t, "rate_limit_reset", decision.reason) + assert.False(t, decision.unfulfillable) + }) + } + + for _, remaining := range []string{"?", "malformed"} { + t.Run("route-only shared "+remaining, func(t *testing.T) { + headers := http.Header{ + "X-Ratelimit-Route-Remaining": {"1"}, + "X-Ratelimit-Auth-Token-Remaining": {"0"}, + "X-Ratelimit-Reset": { + strconv.FormatInt(now.Add(8*time.Second).UnixMilli(), 10), + }, + "X-Ratelimit-Auth-Token-Reset": { + strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10), + }, + } + if remaining == "malformed" { + headers.Set("X-Ratelimit-Global-Remaining", remaining) + } + decision := retryWait( + now, + timePtr(now.Add(5*time.Second)), + 0, + parseRateLimitHeaders(headers), + maximumJitter, + ) + assert.Equal(t, 3*time.Second, decision.delay) + assert.Equal(t, "rate_limit_reset", decision.reason) + assert.False(t, decision.unfulfillable) + }) + } +} + +func TestRetryWaitUsesHTTPDateRetryAfter(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(5 * time.Second) + hints := parseRateLimitHeaders(http.Header{ + "X-Ratelimit-Global-Remaining": {"1"}, + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Reset": { + strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10), + }, + "Retry-After": {now.Add(3 * time.Second).UTC().Format(http.TimeFormat)}, + }) + decision := retryWait( + now, + &deadline, + 0, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.Equal(t, "rate_limit_reset", decision.reason) + assert.False(t, decision.unfulfillable) + assert.Equal(t, 4*time.Second, decision.delay) +} + +func TestRetryWaitFallbackWindows(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(30 * time.Second) + + for retryNumber, want := range map[int]time.Duration{ + 0: 2 * time.Second, + 1: 4 * time.Second, + 2: 8 * time.Second, + 3: 16 * time.Second, + } { + decision := retryWait( + now, + &deadline, + retryNumber, + parsedRateLimitHeaders{}, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.False(t, decision.unfulfillable) + assert.Equal(t, "rate_limit_backoff", decision.reason) + assert.Equal(t, want, decision.delay) + } +} + +func TestRetryWaitRejectsUnrepresentableReset(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + for _, header := range []string{ + "X-Ratelimit-Reset", + "X-Ratelimit-Auth-Token-Reset", + } { + t.Run(header, func(t *testing.T) { + hints := parseRateLimitHeaders(http.Header{ + header: {"9223372036854775808"}, + }) + + decision := retryWait(now, &deadline, 0, hints, uniformJitter) + + assert.True(t, decision.unfulfillable) + }) + } +} + +func TestRetryWaitParsesEachTimingHeader(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + date := now.Add(2 * time.Second).UTC().Format(http.TimeFormat) + tests := []struct { + name string + header string + value string + }{ + {name: "global reset", header: "X-Ratelimit-Reset", value: strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + {name: "auth reset", header: "X-Ratelimit-Auth-Token-Reset", value: strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + {name: "retry after seconds", header: "Retry-After", value: "2"}, + {name: "retry after date", header: "Retry-After", value: date}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + hints := parseRateLimitHeaders(http.Header{tt.header: {tt.value}}) + decision := retryWait( + now, + &deadline, + 0, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.False(t, decision.unfulfillable) + assert.Equal(t, 3*time.Second, decision.delay) + }) + } +} + +func TestRetryWaitHandlesWhitespacePastZeroNegativeAndMalformedHints(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + for _, value := range []string{" 0 ", " ", "-1", "not-a-number", now.Add(-time.Second).Format(http.TimeFormat)} { + hints := parseRateLimitHeaders(http.Header{"Retry-After": {value}}) + decision := retryWait( + now, + &deadline, + 0, + hints, + func(minimum, _ time.Duration) time.Duration { return minimum }, + ) + assert.False(t, decision.unfulfillable, value) + assert.Equal(t, "rate_limit_backoff", decision.reason, value) + assert.Equal(t, time.Second, decision.delay, value) + } +} + +func TestRetryBackoffReturnsRetainedWaitWithoutResampling(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + t.Cleanup(cancel) + jitterCalls := 0 + fetcher := newFlagFetcher(ctx, fetcherOptions{ + jitter: func(_, maximum time.Duration) time.Duration { + jitterCalls++ + return maximum + }, + }) + fetcher.requestLogHook(nil, nil, 0) + + shouldRetry, err := fetcher.checkRetry( + ctx, + &http.Response{StatusCode: http.StatusTooManyRequests, Header: http.Header{}}, + nil, + ) + + require.NoError(t, err) + assert.True(t, shouldRetry) + assert.Equal(t, 1, jitterCalls) + assert.Equal(t, 2*time.Second, fetcher.retryBackoff(0, 0, 0, nil)) + assert.Equal(t, time.Duration(0), fetcher.pendingRetryWait) + assert.Equal(t, time.Duration(0), fetcher.retryBackoff(0, 0, 0, nil)) + assert.Equal(t, 1, jitterCalls) +} + +func TestRetryWaitHandlesMultipleHintsAndOverflow(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + header := http.Header{} + header.Add("X-Ratelimit-Reset", strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)) + header.Add("X-Ratelimit-Reset", strconv.FormatInt(now.Add(4*time.Second).UnixMilli(), 10)) + header.Add("Retry-After", " 1 ") + header.Add("Retry-After", "3") + hints := parseRateLimitHeaders(header) + decision := retryWait( + now, + &deadline, + 0, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.False(t, decision.unfulfillable) + assert.Equal(t, 5*time.Second, decision.delay) + + for _, value := range []string{"9223372037", "18446744073709551616"} { + overflowHints := parseRateLimitHeaders(http.Header{"Retry-After": {value}}) + decision = retryWait(now, &deadline, 0, overflowHints, uniformJitter) + assert.True(t, decision.unfulfillable, value) + } +} + +func TestRetryWaitJitterBounds(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(30 * time.Second) + for name, jitter := range map[string]func(time.Duration, time.Duration) time.Duration{ + "minimum": func(minimum, _ time.Duration) time.Duration { return minimum }, + "maximum": func(_, maximum time.Duration) time.Duration { return maximum }, + } { + t.Run(name, func(t *testing.T) { + decision := retryWait(now, &deadline, 0, parsedRateLimitHeaders{}, jitter) + assert.False(t, decision.unfulfillable) + assert.GreaterOrEqual(t, decision.delay, time.Second) + assert.LessOrEqual(t, decision.delay, 2*time.Second) + }) + } +} + +func TestRetryWaitRejectsDelayAtDeadline(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(2 * time.Second) + hints := parseRateLimitHeaders(http.Header{ + "Retry-After": {"2"}, + }) + + decision := retryWait( + now, + &deadline, + 0, + hints, + func(minimum, _ time.Duration) time.Duration { return minimum }, + ) + + assert.True(t, decision.unfulfillable) +} + +func TestProactiveWaitPairsRemainingAndResetHeaders(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + hints := parseRateLimitHeaders(http.Header{ + "X-Ratelimit-Global-Remaining": {"1"}, + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Auth-Token-Remaining": {"5"}, + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + "X-Ratelimit-Auth-Token-Reset": {strconv.FormatInt(now.Add(8*time.Second).UnixMilli(), 10)}, + }) + + decision := proactiveWait( + now, + &deadline, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + + assert.False(t, decision.unfulfillable) + assert.Equal(t, 3*time.Second, decision.delay) +} + +func TestProactiveWaitRequiresAnExhaustedQuota(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + reset := strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10) + tests := []struct { + name string + header string + resetHeader string + }{ + { + name: "global", + header: "X-Ratelimit-Global-Remaining", + resetHeader: "X-Ratelimit-Reset", + }, + { + name: "route", + header: "X-Ratelimit-Route-Remaining", + resetHeader: "X-Ratelimit-Reset", + }, + { + name: "auth token", + header: "X-Ratelimit-Auth-Token-Remaining", + resetHeader: "X-Ratelimit-Auth-Token-Reset", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + hints := parseRateLimitHeaders(http.Header{ + tt.header: {"0"}, + tt.resetHeader: {reset}, + }) + decision := proactiveWait( + now, + &deadline, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.False(t, decision.unfulfillable) + assert.Equal(t, 3*time.Second, decision.delay) + }) + } +} + +func TestProactiveWaitIgnoresPositiveAndUnknownRemaining(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + tests := []struct { + name string + remaining string + remainingHeader string + resetHeader string + }{ + { + name: "positive global", + remaining: "1", + remainingHeader: "X-Ratelimit-Global-Remaining", + resetHeader: "X-Ratelimit-Reset", + }, + { + name: "unknown route", + remaining: "unknown", + remainingHeader: "X-Ratelimit-Route-Remaining", + resetHeader: "X-Ratelimit-Reset", + }, + { + name: "positive token", + remaining: "1", + remainingHeader: "X-Ratelimit-Auth-Token-Remaining", + resetHeader: "X-Ratelimit-Auth-Token-Reset", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for _, remaining := range []string{tt.remaining, "unknown"} { + reset := strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10) + hints := parseRateLimitHeaders(http.Header{ + tt.remainingHeader: {remaining}, + tt.resetHeader: {reset}, + }) + decision := proactiveWait( + now, + &deadline, + hints, + uniformJitter, + ) + assert.False(t, decision.unfulfillable, remaining) + assert.Zero(t, decision.delay, remaining) + } + }) + } +} + +func TestProactiveWaitUsesFallbackForMissingOrMalformedReset(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + tests := []struct { + name string + remaining string + remainingHeader string + resetHeader string + }{ + { + name: "global missing", + remaining: "0", + remainingHeader: "X-Ratelimit-Global-Remaining", + resetHeader: "X-Ratelimit-Reset", + }, + { + name: "route malformed", + remaining: "0", + remainingHeader: "X-Ratelimit-Route-Remaining", + resetHeader: "X-Ratelimit-Reset", + }, + { + name: "token missing", + remaining: "0", + remainingHeader: "X-Ratelimit-Auth-Token-Remaining", + resetHeader: "X-Ratelimit-Auth-Token-Reset", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + headers := http.Header{tt.remainingHeader: {tt.remaining}} + if tt.name == "route malformed" { + headers.Set(tt.resetHeader, "not-a-reset") + } + hints := parseRateLimitHeaders(headers) + decision := proactiveWait( + now, + &deadline, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.False(t, decision.unfulfillable) + assert.Equal(t, 2*time.Second, decision.delay) + }) + } +} + +func TestProactiveWaitHonorsRetryAfterMinimum(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + tests := []struct { + name string + retryAfter []string + wantDelay time.Duration + unfulfillable bool + }{ + {name: "seconds", retryAfter: []string{"5"}, wantDelay: 6 * time.Second}, + {name: "multiple values", retryAfter: []string{"3", "5"}, wantDelay: 6 * time.Second}, + { + name: "HTTP date", + retryAfter: []string{now.Add(5 * time.Second).UTC().Format(http.TimeFormat)}, + wantDelay: 6 * time.Second, + }, + {name: "duration overflow", retryAfter: []string{"9223372037"}, unfulfillable: true}, + { + name: "integer overflow", + retryAfter: []string{"18446744073709551616"}, + unfulfillable: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + hints := parseRateLimitHeaders(http.Header{ + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Reset": { + strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10), + }, + "Retry-After": tt.retryAfter, + }) + decision := proactiveWait( + now, + &deadline, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + assert.Equal(t, tt.unfulfillable, decision.unfulfillable) + if tt.unfulfillable { + assert.Zero(t, decision.delay) + } else { + assert.Equal(t, tt.wantDelay, decision.delay) + } + }) + } +} + +func TestProactiveWaitDoesNotPairTokenResetWithAnotherQuota(t *testing.T) { + now := time.Unix(1_800_000_000, 0) + deadline := now.Add(10 * time.Second) + hints := parseRateLimitHeaders(http.Header{ + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Global-Remaining": {"1"}, + "X-Ratelimit-Auth-Token-Remaining": {"1"}, + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + "X-Ratelimit-Auth-Token-Reset": {strconv.FormatInt(now.Add(8*time.Second).UnixMilli(), 10)}, + }) + + decision := proactiveWait( + now, + &deadline, + hints, + func(_, maximum time.Duration) time.Duration { return maximum }, + ) + + assert.False(t, decision.unfulfillable) + assert.Equal(t, 3*time.Second, decision.delay) +} + +func TestFetcherRetriesOnlyRateLimitedPages(t *testing.T) { + config := serveFlagPages(t, false, []flagPageResponse{ + {body: `{"items":[{"key":"first"}]}`}, + {offset: 100, status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {offset: 100, status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {offset: 100, body: `{"items":[{"key":"second"}]}`}, + {offset: 200, body: `{"items":[]}`}, + }) + fetcher, _ := newTestFetcher(t) + waits := []time.Duration{} + setTestBackoff(fetcher, func(wait time.Duration) time.Duration { + waits = append(waits, wait) + return 0 + }) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.NoError(t, err) + assert.Equal(t, []string{"first", "second"}, keys(flags)) + assert.Equal(t, 5, fetcher.diagnostics.httpAttempts) + assert.Equal(t, 2, fetcher.diagnostics.retries) + assert.Equal(t, 2, fetcher.diagnostics.responses429) + assert.Equal(t, []time.Duration{2 * time.Second, 4 * time.Second}, waits) +} + +func TestFetcherRetriesAfterExhaustedQuotaReset(t *testing.T) { + testDeadline := time.Now().Add(30 * time.Second).Truncate(time.Millisecond) + parentCtx, cancel := context.WithDeadline(context.Background(), testDeadline) + t.Cleanup(cancel) + fetcher, logs := newTestFetcherWithContext(t, parentCtx) + deadline, ok := fetcher.ctx.Deadline() + require.True(t, ok) + now := deadline.Add(-5 * time.Second) + fetcher.clock = func() time.Time { return now } + + config := serveFlagPages(t, false, []flagPageResponse{ + { + status: http.StatusTooManyRequests, + headers: http.Header{ + "X-Ratelimit-Global-Remaining": {"1"}, + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Auth-Token-Remaining": {"5"}, + "X-Ratelimit-Reset": { + strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10), + }, + "X-Ratelimit-Auth-Token-Reset": { + strconv.FormatInt(now.Add(8*time.Second).UnixMilli(), 10), + }, + }, + body: "{\"message\":\"rate limited\"}", + }, + {body: "{\"items\":[]}"}, + }) + + var waits []time.Duration + setTestBackoff(fetcher, func(wait time.Duration) time.Duration { + waits = append(waits, wait) + return 0 + }) + fetcher.diagnostics.emitStart() + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + fetcher.diagnostics.emitSummary(fetcher.clock(), fetcher.outcome(err), err == nil) + + require.NoError(t, err) + assert.Empty(t, flags) + assert.Equal(t, 2, fetcher.diagnostics.httpAttempts) + assert.Equal(t, 1, fetcher.diagnostics.retries) + assert.Equal(t, 1, fetcher.diagnostics.responses429) + assert.Equal(t, 1, fetcher.diagnostics.retryWaits) + assert.Equal(t, []time.Duration{3 * time.Second}, waits) + + records := diagnosticRecords(t, *logs) + var summary map[string]any + var wait map[string]any + for _, record := range records { + switch record["event"] { + case "summary": + summary = record + case "wait": + wait = record + } + } + require.NotNil(t, wait) + assert.Equal(t, "rate_limit_reset", wait["reason"]) + assert.Equal(t, float64(3000), wait["intended_wait_ms"]) + require.NotNil(t, summary) + assert.Equal(t, "success", summary["outcome"]) + assert.Equal(t, true, summary["inventory_complete"]) +} + +func TestFetcherRetriesAtEveryPaginationPosition(t *testing.T) { + tests := []struct { + name string + includeArchived bool + pages []flagPageResponse + wantKeys []string + }{ + { + name: "first active", + pages: []flagPageResponse{ + {status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {body: `{"items":[{"key":"active-first"}]}`}, + {offset: 100, body: `{"items":[]}`}, + }, + wantKeys: []string{"active-first"}, + }, + { + name: "middle active", + pages: []flagPageResponse{ + {body: `{"items":[{"key":"active-first"}]}`}, + {offset: 100, status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {offset: 100, body: `{"items":[{"key":"active-middle"}]}`}, + {offset: 200, body: `{"items":[]}`}, + }, + wantKeys: []string{"active-first", "active-middle"}, + }, + { + name: "terminal empty page", + pages: []flagPageResponse{ + {body: `{"items":[{"key":"active-first"}]}`}, + {offset: 100, status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {offset: 100, body: `{"items":[]}`}, + }, + wantKeys: []string{"active-first"}, + }, + { + name: "first archived", + includeArchived: true, + pages: []flagPageResponse{ + {body: `{"items":[]}`}, + {filter: "state:archived", status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {filter: "state:archived", body: `{"items":[{"key":"archived-first"}]}`}, + {filter: "state:archived", offset: 100, body: `{"items":[]}`}, + }, + wantKeys: []string{"archived-first"}, + }, + { + name: "middle archived", + includeArchived: true, + pages: []flagPageResponse{ + {body: `{"items":[]}`}, + {filter: "state:archived", body: `{"items":[{"key":"archived-first"}]}`}, + {filter: "state:archived", offset: 100, status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {filter: "state:archived", offset: 100, body: `{"items":[{"key":"archived-middle"}]}`}, + {filter: "state:archived", offset: 200, body: `{"items":[]}`}, + }, + wantKeys: []string{"archived-first", "archived-middle"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + config := serveFlagPages(t, tt.includeArchived, tt.pages) + fetcher, _ := newTestFetcher(t) + params := urlValuesForFlags() + flags, err := fetcher.getFlags(config, params, flagCollectionActive, tt.includeArchived) + require.NoError(t, err) + if tt.includeArchived { + params.Set("filter", "state:archived") + archived, archivedErr := fetcher.getFlags(config, params, flagCollectionArchived, false) + require.NoError(t, archivedErr) + flags = append(flags, archived...) + } + assert.Equal(t, tt.wantKeys, keys(flags)) + }) + } +} + +func TestFetcherDoesNotRetryNon429Statuses(t *testing.T) { + for _, status := range []int{ + http.StatusBadRequest, + http.StatusUnauthorized, + http.StatusForbidden, + http.StatusNotFound, + http.StatusInternalServerError, + http.StatusNotImplemented, + http.StatusBadGateway, + http.StatusServiceUnavailable, + } { + t.Run(strconv.Itoa(status), func(t *testing.T) { + config := serveFlagPages(t, false, []flagPageResponse{{status: status, body: `{"message":"sentinel"}`}}) + fetcher, _ := newTestFetcher(t) + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) + assert.Contains(t, err.Error(), "status="+strconv.Itoa(status)) + }) + } +} + +func TestFetcherDoesNotRetryMalformedSuccessJSON(t *testing.T) { + config := serveFlagPages(t, false, []flagPageResponse{{status: http.StatusOK, body: "not JSON"}}) + fetcher, _ := newTestFetcher(t) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) + assert.Equal(t, "decode_error", fetcher.outcome(err)) + var syntaxErr *json.SyntaxError + assert.ErrorAs(t, err, &syntaxErr) +} + +func TestFetcherClassifiesBodyReadErrors(t *testing.T) { + readSentinel := errors.New("body-read-sentinel") + canceledReadErr := fmt.Errorf("body-read-secret: %w", context.Canceled) + deadlineReadErr := fmt.Errorf("body-read-secret: %w", context.DeadlineExceeded) + timeoutErr := &net.OpError{ + Op: "read", + Net: "tcp", + Err: os.ErrDeadlineExceeded, + } + timeoutReadErr := fmt.Errorf("body-read-secret: %w", timeoutErr) + canceledDuringReadErr := errors.New("body-read-cancel-sentinel") + tests := []struct { + name string + readErr error + wantReason fetchErrorReason + wantOutcome string + wantContext error + sentinel string + cancelDuringRead bool + wantTimeout bool + }{ + { + name: "wrapped cancellation", + readErr: canceledReadErr, + wantReason: reasonCanceled, + wantOutcome: "canceled", + wantContext: context.Canceled, + sentinel: "body-read-secret", + }, + { + name: "wrapped deadline", + readErr: deadlineReadErr, + wantReason: reasonDeadline, + wantOutcome: "deadline", + wantContext: context.DeadlineExceeded, + sentinel: "body-read-secret", + }, + { + name: "timeout net error", + readErr: timeoutReadErr, + wantReason: reasonDeadline, + wantOutcome: "deadline", + sentinel: "body-read-secret", + wantTimeout: true, + }, + { + name: "non-timeout read error", + readErr: readSentinel, + wantReason: reasonDecode, + wantOutcome: "decode_error", + sentinel: "body-read-sentinel", + }, + { + name: "context canceled during read", + readErr: canceledDuringReadErr, + wantReason: reasonCanceled, + wantOutcome: "canceled", + wantContext: context.Canceled, + sentinel: "body-read-cancel-sentinel", + cancelDuringRead: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + parentCtx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + var closed atomic.Int32 + var reader io.Reader = iotest.ErrReader(tt.readErr) + if tt.cancelDuringRead { + reader = readFunc(func([]byte) (int, error) { + cancel() + return 0, tt.readErr + }) + } + transport := roundTripFunc(func(request *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Header: http.Header{}, + Body: &trackingBody{ + Reader: reader, + closed: &closed, + }, + Request: request, + }, nil + }) + fetcher, logs := newTestFetcherWithContext(t, parentCtx) + fetcher.client.HTTPClient.Transport = transport + fetcher.diagnostics.emitStart() + + flags, err := fetcher.getFlags( + &lcr.Config{ + LdInstance: "https://example.test", + LdProject: "project", + LdEnvironment: "production", + ApiToken: "token-sentinel", + }, + urlValuesForFlags(), + flagCollectionActive, + false, + ) + fetcher.diagnostics.emitSummary(fetcher.clock(), fetcher.outcome(err), err == nil) + + require.Error(t, err) + assert.Empty(t, flags) + var fetchErr *flagFetchError + require.ErrorAs(t, err, &fetchErr) + assert.Equal(t, http.StatusOK, fetchErr.status) + assert.Equal(t, tt.wantReason, fetchErr.reason) + assert.Equal(t, tt.wantOutcome, fetcher.outcome(err)) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) + assert.Equal(t, 0, fetcher.diagnostics.retries) + assert.Equal(t, 0, fetcher.diagnostics.successfulPages) + assert.Equal(t, int32(1), closed.Load()) + assert.ErrorIs(t, err, tt.readErr) + if tt.wantContext != nil { + assert.ErrorIs(t, err, tt.wantContext) + } + if tt.wantTimeout { + var timeout net.Error + require.ErrorAs(t, err, &timeout) + assert.True(t, timeout.Timeout()) + } + assert.NotContains(t, err.Error(), tt.sentinel) + + records := diagnosticRecords(t, *logs) + var summary map[string]any + for _, record := range records { + if record["event"] == "summary" { + summary = record + } + assert.NotContains(t, fmt.Sprint(record), tt.sentinel) + } + require.NotNil(t, summary) + assert.Equal(t, false, summary["inventory_complete"]) + assert.Equal(t, float64(http.StatusOK), summary["last_http_status"]) + assert.Equal(t, flagCollectionActive, summary["failed_collection"]) + assert.Equal(t, float64(0), summary["failed_offset"]) + }) + } +} + +func TestFetcherReturnsNoPartialInventoryAfterLaterPageReadCancellation(t *testing.T) { + var closed atomic.Int32 + readErr := fmt.Errorf("later-page-secret: %w", context.Canceled) + var requestNumber atomic.Int32 + transport := roundTripFunc(func(request *http.Request) (*http.Response, error) { + index := requestNumber.Add(1) + var reader io.Reader = strings.NewReader("{\"items\":[{\"key\":\"partial\"}]}") + if index == 2 { + reader = iotest.ErrReader(readErr) + } + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Header: http.Header{}, + Body: &trackingBody{ + Reader: reader, + closed: &closed, + }, + Request: request, + }, nil + }) + fetcher, logs := newTestFetcher(t) + fetcher.client.HTTPClient.Transport = transport + fetcher.diagnostics.emitStart() + + flags, err := fetcher.getFlags( + &lcr.Config{ + LdInstance: "https://example.test", + LdProject: "project", + LdEnvironment: "production", + ApiToken: "token-sentinel", + }, + urlValuesForFlags(), + flagCollectionActive, + false, + ) + fetcher.diagnostics.emitSummary(fetcher.clock(), fetcher.outcome(err), err == nil) + + require.Error(t, err) + assert.Empty(t, flags) + assert.ErrorIs(t, err, readErr) + assert.Equal(t, "canceled", fetcher.outcome(err)) + assert.Equal(t, 2, fetcher.diagnostics.httpAttempts) + assert.Equal(t, 1, fetcher.diagnostics.successfulPages) + assert.Equal(t, int32(2), closed.Load()) + var fetchErr *flagFetchError + require.ErrorAs(t, err, &fetchErr) + assert.Equal(t, http.StatusOK, fetchErr.status) + assert.Equal(t, 100, fetchErr.offset) + + records := diagnosticRecords(t, *logs) + var summary map[string]any + for _, record := range records { + if record["event"] == "summary" { + summary = record + } + assert.NotContains(t, fmt.Sprint(record), "later-page-secret") + } + require.NotNil(t, summary) + assert.Equal(t, false, summary["inventory_complete"]) + assert.Equal(t, float64(http.StatusOK), summary["last_http_status"]) + assert.Equal(t, flagCollectionActive, summary["failed_collection"]) + assert.Equal(t, float64(100), summary["failed_offset"]) +} + +func TestFetcherDoesNotRetryTransportErrors(t *testing.T) { + transport := roundTripFunc(func(*http.Request) (*http.Response, error) { + return nil, errors.New("transport sentinel") + }) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + fetcher := newFlagFetcher(ctx, fetcherOptions{httpClient: &http.Client{Transport: transport}}) + config := &lcr.Config{LdInstance: "https://example.test", LdProject: "project", LdEnvironment: "production", ApiToken: "token-sentinel"} + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) + assert.Equal(t, "transport_error", fetcher.outcome(err)) + assert.NotContains(t, err.Error(), "example.test") + assert.NotContains(t, err.Error(), "token-sentinel") +} + +func TestFetcherRateLimitExhaustionReturnsEmptyInventory(t *testing.T) { + pages := make([]flagPageResponse, maxRetries+1) + for i := range pages { + pages[i] = flagPageResponse{ + offset: 0, + status: http.StatusTooManyRequests, + body: `{"message":"rate limited"}`, + } + } + config := serveFlagPages(t, false, pages) + fetcher, _ := newTestFetcher(t) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, 5, fetcher.diagnostics.httpAttempts) + assert.Equal(t, 4, fetcher.diagnostics.retries) + assert.Equal(t, 5, fetcher.diagnostics.responses429) + assert.Equal(t, "rate_limit_exhausted", fetcher.outcome(err)) +} + +func TestFetcherDoesNotRetryPermanentHTTPStatus(t *testing.T) { + config := serveFlagPages(t, false, []flagPageResponse{ + {status: http.StatusServiceUnavailable, body: `{"message":"unavailable"}`}, + }) + fetcher, _ := newTestFetcher(t) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) + assert.Equal(t, "status=503 reason=http_error", errorSuffix(err)) +} + +func TestFetcherAppliesProactiveWaitBeforeNextPage(t *testing.T) { + now := time.Now().Truncate(time.Second) + waits := []time.Duration{} + config := serveFlagPages(t, false, []flagPageResponse{ + { + body: `{"items":[{"key":"first"}]}`, + headers: http.Header{ + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + }, + }, + {offset: 100, body: `{"items":[]}`}, + }) + fetcher, _ := newTestFetcher(t) + fetcher.clock = func() time.Time { return now } + fetcher.wait = func(_ context.Context, wait time.Duration) error { + waits = append(waits, wait) + now = now.Add(wait) + return nil + } + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.NoError(t, err) + assert.Len(t, flags, 1) + assert.Equal(t, []time.Duration{3 * time.Second}, waits) + assert.Equal(t, 1, fetcher.diagnostics.proactiveWaits) + assert.Equal(t, int64(3000), fetcher.diagnostics.scheduledWait.Milliseconds()) +} + +func TestFetcherWaitsAtActiveArchivedBoundary(t *testing.T) { + now := time.Now().Truncate(time.Second) + waits := 0 + config := serveFlagPages(t, true, []flagPageResponse{ + { + filter: "", + body: `{"items":[]}`, + headers: http.Header{ + "X-Ratelimit-Global-Remaining": {"0"}, + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + }, + }, + {filter: "state:archived", body: `{"items":[]}`}, + }) + fetcher, _ := newTestFetcher(t) + fetcher.clock = func() time.Time { return now } + fetcher.wait = func(_ context.Context, wait time.Duration) error { + waits++ + now = now.Add(wait) + return nil + } + + active, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, true) + require.NoError(t, err) + params := urlValuesForFlags() + params.Set("filter", "state:archived") + archived, err := fetcher.getFlags(config, params, flagCollectionArchived, false) + + require.NoError(t, err) + assert.Empty(t, active) + assert.Empty(t, archived) + assert.Equal(t, 1, waits) +} + +func TestFetcherDoesNotWaitAfterFinalArchivedPage(t *testing.T) { + now := time.Now().Truncate(time.Second) + waits := 0 + config := serveFlagPages(t, true, []flagPageResponse{ + {filter: "state:archived", body: `{"items":[]}`, headers: http.Header{ + "X-Ratelimit-Global-Remaining": {"0"}, + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + }}, + }) + fetcher, _ := newTestFetcher(t) + fetcher.clock = func() time.Time { return now } + fetcher.wait = func(_ context.Context, wait time.Duration) error { + waits++ + return nil + } + + params := urlValuesForFlags() + params.Set("filter", "state:archived") + flags, err := fetcher.getFlags(config, params, flagCollectionArchived, false) + + require.NoError(t, err) + assert.Empty(t, flags) + assert.Equal(t, 0, waits) + assert.Equal(t, 0, fetcher.diagnostics.proactiveWaits) +} + +func TestFetcherFailsBeforeArchivedRequestWhenResetCannotFit(t *testing.T) { + now := time.Now().Truncate(time.Second) + config := serveFlagPages(t, true, []flagPageResponse{ + { + body: `{"items":[]}`, + headers: http.Header{ + "X-Ratelimit-Global-Remaining": {"0"}, + "X-Ratelimit-Reset": { + strconv.FormatInt(now.Add(60*time.Second).UnixMilli(), 10), + }, + }, + }, + }) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + t.Cleanup(cancel) + fetcher, _ := newTestFetcherWithContext(t, ctx) + fetcher.clock = func() time.Time { return now } + + active, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, true) + require.NoError(t, err) + assert.Empty(t, active) + + params := urlValuesForFlags() + params.Set("filter", "state:archived") + archived, err := fetcher.getFlags(config, params, flagCollectionArchived, false) + + require.Error(t, err) + assert.Empty(t, archived) + assert.Equal(t, "rate_limit_deadline", fetcher.outcome(err)) + assert.Contains(t, err.Error(), "collection=archived offset=0") + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) +} + +func TestFetcherCancellationDuringRetryWait(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + config := serveFlagPages(t, false, []flagPageResponse{ + {status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + }) + fetcher, _ := newTestFetcherWithContext(t, ctx) + setTestBackoff(fetcher, func(wait time.Duration) time.Duration { + cancel() + return wait + }) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, "canceled", fetcher.outcome(err)) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) +} + +func TestFetcherCancellationDuringProactiveWait(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + now := time.Now().Truncate(time.Second) + config := serveFlagPages(t, false, []flagPageResponse{ + { + body: `{"items":[{"key":"first"}]}`, + headers: http.Header{ + "X-Ratelimit-Route-Remaining": {"0"}, + "X-Ratelimit-Reset": {strconv.FormatInt(now.Add(2*time.Second).UnixMilli(), 10)}, + }, + }, + }) + fetcher, _ := newTestFetcherWithContext(t, ctx) + fetcher.clock = func() time.Time { return now } + fetcher.wait = func(waitContext context.Context, _ time.Duration) error { + cancel() + return waitContext.Err() + } + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, "canceled", fetcher.outcome(err)) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) +} + +func TestCallerDeadlineWinsRateLimitRetryDeadline(t *testing.T) { + config := serveFlagPages(t, false, []flagPageResponse{{ + status: http.StatusTooManyRequests, + headers: http.Header{ + "Retry-After": {strconv.FormatInt(60, 10)}, + }, + body: `{"message":"rate limited"}`, + }}) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + t.Cleanup(cancel) + + flags, err := GetAllFlags(ctx, config) + + require.Error(t, err) + assert.Empty(t, flags) + assert.ErrorIs(t, err, context.DeadlineExceeded) + assert.Contains(t, err.Error(), "reason=rate_limit_deadline") +} + +func TestCallerDeadlineCoversActiveAndArchivedCollections(t *testing.T) { + var requests atomic.Int32 + startedArchived := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Query().Get("filter") == "state:archived" { + requests.Add(1) + close(startedArchived) + <-r.Context().Done() + return + } + requests.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"items":[]}`)) + })) + t.Cleanup(server.Close) + config := &lcr.Config{ + LdInstance: server.URL, + LdProject: "test-project", + LdEnvironment: "production", + ApiToken: "token-sentinel", + IncludeArchivedFlags: true, + } + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + t.Cleanup(cancel) + + flags, err := GetAllFlags(ctx, config) + + require.Error(t, err) + assert.Empty(t, flags) + assert.ErrorIs(t, err, context.DeadlineExceeded) + assert.Equal(t, int32(2), requests.Load()) + select { + case <-startedArchived: + default: + t.Fatal("archived collection was not requested") + } +} + +func TestFetcherCancellationDuringHTTP(t *testing.T) { + started := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + close(started) + <-r.Context().Done() + })) + t.Cleanup(server.Close) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + fetcher, _ := newTestFetcherWithContext(t, ctx) + config := &lcr.Config{ + LdInstance: server.URL, + LdProject: "test-project", + LdEnvironment: "production", + ApiToken: "token-sentinel", + } + result := make(chan error, 1) + go func() { + _, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + result <- err + }() + <-started + cancel() + err := <-result + + require.Error(t, err) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, "canceled", fetcher.outcome(err)) + assert.Equal(t, 1, fetcher.diagnostics.httpAttempts) +} + +func TestFetcherCancellationBeforeFirstRequest(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + config := serveFlagPages(t, false, nil) + fetcher, _ := newTestFetcherWithContext(t, ctx) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.ErrorIs(t, err, context.Canceled) + assert.Equal(t, "canceled", fetcher.outcome(err)) + assert.Equal(t, 0, fetcher.diagnostics.httpAttempts) +} + +func TestFetcherResponseBodiesCloseOnSuccessAndRetry(t *testing.T) { + var closed atomic.Int32 + responses := []trackingResponse{ + {status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {status: http.StatusOK, body: `{"items":[]}`}, + } + transport := &trackingRoundTripper{responses: responses, closed: &closed} + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + fetcher := newFlagFetcher(ctx, fetcherOptions{ + httpClient: &http.Client{Transport: transport}, + jitter: func(_, maximum time.Duration) time.Duration { return maximum }, + }) + setTestBackoff(fetcher, func(time.Duration) time.Duration { return 0 }) + config := &lcr.Config{LdInstance: "https://example.test", LdProject: "project", LdEnvironment: "production", ApiToken: "sentinel-token"} + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.NoError(t, err) + assert.Empty(t, flags) + assert.Equal(t, int32(2), closed.Load()) +} + +func TestFetcherClosesFinalResponseOnDecodeError(t *testing.T) { + var closed atomic.Int32 + transport := &trackingRoundTripper{ + responses: []trackingResponse{{status: http.StatusOK, body: "not JSON"}}, + closed: &closed, + } + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + fetcher := newFlagFetcher(ctx, fetcherOptions{ + httpClient: &http.Client{Transport: transport}, + }) + config := &lcr.Config{ + LdInstance: "https://example.test", + LdProject: "project", + LdEnvironment: "production", + ApiToken: "sentinel-token", + } + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, int32(1), closed.Load()) +} + +func TestFetcherClosesFinalResponseReturnedWithError(t *testing.T) { + var closed atomic.Int32 + transport := roundTripFunc(func(request *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusTooManyRequests, + Status: "429 Too Many Requests", + Header: http.Header{"Retry-After": {"60"}}, + Body: &trackingBody{ + Reader: strings.NewReader(`{"message":"rate limited"}`), + closed: &closed, + }, + Request: request, + }, nil + }) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + t.Cleanup(cancel) + fetcher := newFlagFetcher(ctx, fetcherOptions{ + httpClient: &http.Client{Transport: transport}, + }) + config := &lcr.Config{ + LdInstance: "https://example.test", + LdProject: "project", + LdEnvironment: "production", + ApiToken: "sentinel-token", + } + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.Equal(t, "rate_limit_deadline", fetcher.outcome(err)) + assert.Equal(t, int32(1), closed.Load()) +} + +func TestFetchErrorsAndDiagnosticsDoNotEchoSensitiveValues(t *testing.T) { + const sentinel = "flag-secret-name-authorization-token" + config := serveFlagPages(t, false, []flagPageResponse{ + {status: http.StatusBadGateway, body: `{"message":"` + sentinel + `"}`}, + }) + fetcher, logs := newTestFetcher(t) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.Error(t, err) + assert.Empty(t, flags) + assert.NotContains(t, err.Error(), sentinel) + assert.NotContains(t, strings.Join(*logs, "\n"), sentinel) +} + +func TestMalformedRateLimitHeadersAreNotLogged(t *testing.T) { + const sentinel = "malformed-header-secret" + config := serveFlagPages(t, false, []flagPageResponse{ + { + status: http.StatusTooManyRequests, + body: `{"message":"rate limited"}`, + headers: http.Header{ + "Retry-After": {sentinel}, + "X-Ratelimit-Reset": {sentinel}, + }, + }, + {body: `{"items":[]}`}, + }) + fetcher, logs := newTestFetcher(t) + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + + require.NoError(t, err) + assert.Empty(t, flags) + assert.NotContains(t, strings.Join(*logs, "\n"), sentinel) +} + +func TestFetcherDiagnosticsHaveOneStartAndSummary(t *testing.T) { + tests := []struct { + name string + pages []flagPageResponse + wantAttempts float64 + wantRetries float64 + want429 float64 + wantSuccessful float64 + wantNonempty float64 + wantRetryWaits float64 + wantProactive float64 + wantWaitEvents int + wantRateLimitLog int + }{ + { + name: "healthy", + pages: []flagPageResponse{ + {body: `{"items":[{"key":"healthy"}]}`}, + {offset: 100, body: `{"items":[]}`}, + }, + wantAttempts: 2, + wantSuccessful: 2, + wantNonempty: 1, + }, + { + name: "recovered rate limit", + pages: []flagPageResponse{ + {status: http.StatusTooManyRequests, body: `{"message":"rate limited"}`}, + {body: `{"items":[{"key":"recovered"}]}`}, + {offset: 100, body: `{"items":[]}`}, + }, + wantAttempts: 3, + wantRetries: 1, + want429: 1, + wantSuccessful: 2, + wantNonempty: 1, + wantRetryWaits: 1, + wantWaitEvents: 1, + wantRateLimitLog: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + config := serveFlagPages(t, false, tt.pages) + fetcher, logs := newTestFetcher(t) + fetcher.diagnostics.emitStart() + + _, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + fetcher.diagnostics.emitSummary(fetcher.clock(), fetcher.outcome(err), err == nil) + + records := diagnosticRecords(t, *logs) + counts := map[string]int{} + var summary map[string]any + for _, record := range records { + event, ok := record["event"].(string) + require.True(t, ok) + counts[event]++ + if event == "summary" { + summary = record + } + } + + assert.Equal(t, 1, counts["start"]) + assert.Equal(t, 1, counts["summary"]) + assert.Equal(t, tt.wantWaitEvents, counts["wait"]) + assert.Equal(t, tt.wantRateLimitLog, counts["rate_limit"]) + require.NotNil(t, summary) + assert.Equal(t, "success", summary["outcome"]) + assert.Equal(t, true, summary["inventory_complete"]) + assert.Equal(t, tt.wantAttempts, summary["http_attempts"]) + assert.Equal(t, tt.wantRetries, summary["retries"]) + assert.Equal(t, tt.want429, summary["responses_429"]) + assert.Equal(t, tt.wantSuccessful, summary["successful_pages"]) + assert.Equal(t, tt.wantNonempty, summary["nonempty_pages"]) + assert.Equal(t, tt.wantRetryWaits, summary["retry_waits"]) + assert.Equal(t, tt.wantProactive, summary["proactive_waits"]) + }) + } +} + +func TestFetcherDiagnosticsReportRateLimitFailure(t *testing.T) { + pages := make([]flagPageResponse, maxRetries+1) + for index := range pages { + pages[index] = flagPageResponse{ + status: http.StatusTooManyRequests, + body: `{"message":"rate limited"}`, + } + } + config := serveFlagPages(t, false, pages) + fetcher, logs := newTestFetcher(t) + fetcher.diagnostics.emitStart() + + flags, err := fetcher.getFlags(config, urlValuesForFlags(), flagCollectionActive, false) + fetcher.diagnostics.emitSummary(fetcher.clock(), fetcher.outcome(err), err == nil) + + require.Error(t, err) + assert.Empty(t, flags) + records := diagnosticRecords(t, *logs) + counts := map[string]int{} + var summary map[string]any + for _, record := range records { + event, ok := record["event"].(string) + require.True(t, ok) + counts[event]++ + if event == "summary" { + summary = record + } + } + + assert.Equal(t, 1, counts["start"]) + assert.Equal(t, 1, counts["summary"]) + assert.Equal(t, maxRetries, counts["wait"]) + assert.Equal(t, maxRetries+1, counts["rate_limit"]) + require.NotNil(t, summary) + assert.Equal(t, "rate_limit_exhausted", summary["outcome"]) + assert.Equal(t, false, summary["inventory_complete"]) + assert.Equal(t, float64(maxRetries+1), summary["http_attempts"]) + assert.Equal(t, float64(maxRetries), summary["retries"]) + assert.Equal(t, float64(maxRetries+1), summary["responses_429"]) + assert.Equal(t, float64(0), summary["successful_pages"]) + assert.Equal(t, float64(0), summary["nonempty_pages"]) + assert.Equal(t, "active", summary["failed_collection"]) + assert.Equal(t, float64(0), summary["failed_offset"]) +} + +func diagnosticRecords(t *testing.T, logs []string) []map[string]any { + t.Helper() + records := make([]map[string]any, 0, len(logs)) + for _, logLine := range logs { + const prefix = "LD_FLAG_FETCH " + require.True(t, strings.HasPrefix(logLine, prefix)) + var record map[string]any + require.NoError(t, json.Unmarshal([]byte(strings.TrimSpace(strings.TrimPrefix(logLine, prefix))), &record)) + records = append(records, record) + } + return records +} + +func newTestFetcher(t *testing.T) (*flagFetcher, *[]string) { + return newTestFetcherWithContext(t, context.Background()) +} + +func newTestFetcherWithContext(t *testing.T, parent context.Context) (*flagFetcher, *[]string) { + ctx, cancel := context.WithTimeout(parent, 30*time.Second) + t.Cleanup(cancel) + logs := &[]string{} + fetcher := newFlagFetcher(ctx, fetcherOptions{ + jitter: func(_, maximum time.Duration) time.Duration { return maximum }, + log: func(format string, args ...any) { + *logs = append(*logs, fmt.Sprintf(format, args...)) + }, + }) + setTestBackoff(fetcher, func(time.Duration) time.Duration { return 0 }) + return fetcher, logs +} + +func setTestBackoff(fetcher *flagFetcher, transform func(time.Duration) time.Duration) { + fetcher.client.Backoff = func(minimum, maximum time.Duration, attempt int, response *http.Response) time.Duration { + return transform(fetcher.retryBackoff(minimum, maximum, attempt, response)) + } +} + +func timePtr(value time.Time) *time.Time { + return &value +} + +func urlValuesForFlags() url.Values { + return url.Values{"env": {"production"}, "limit": {"100"}} +} + +func keys(flags []ldapi.FeatureFlag) []string { + keys := make([]string, len(flags)) + for i, flag := range flags { + keys[i] = flag.Key + } + return keys +} + +func errorSuffix(err error) string { + message := err.Error() + if index := strings.Index(message, "status="); index >= 0 { + return message[index:] + } + return message +} + +type trackingResponse struct { + status int + body string + headers http.Header +} + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return f(request) +} + +type readFunc func([]byte) (int, error) + +func (f readFunc) Read(buffer []byte) (int, error) { + return f(buffer) +} + +type trackingRoundTripper struct { + responses []trackingResponse + closed *atomic.Int32 + index atomic.Int32 +} + +func (t *trackingRoundTripper) RoundTrip(_ *http.Request) (*http.Response, error) { + index := int(t.index.Add(1)) - 1 + response := t.responses[index] + status := response.status + if status == 0 { + status = http.StatusOK + } + headers := response.headers + if headers == nil { + headers = http.Header{} + } + return &http.Response{ + StatusCode: status, + Status: strconv.Itoa(status), + Header: headers, + Body: &trackingBody{ + Reader: strings.NewReader(response.body), + closed: t.closed, + }, + Request: &http.Request{}, + }, nil +} + +type trackingBody struct { + io.Reader + closed *atomic.Int32 +} + +func (b *trackingBody) Close() error { + b.closed.Add(1) + return nil +} + +var _ io.ReadCloser = (*trackingBody)(nil) diff --git a/main.go b/main.go index dffc8b30..c1dccf32 100644 --- a/main.go +++ b/main.go @@ -38,7 +38,7 @@ func main() { failExit(err) } - flags, err := ldclient.GetAllFlags(config) + flags, err := ldclient.GetAllFlags(ctx, config) failExit(err) if len(flags) == 0 {