diff --git a/internal/adapters/distributor/client.go b/internal/adapters/distributor/client.go index 751e631..caca28e 100644 --- a/internal/adapters/distributor/client.go +++ b/internal/adapters/distributor/client.go @@ -24,6 +24,7 @@ type Client struct { TokenEnv string Timeout time.Duration newUploadClient uploadClientFactory + pollWait func(context.Context, time.Duration) error } type UploadRequest struct { @@ -126,6 +127,7 @@ func newClient(cfg config.DistributorNotifyConfig, factory uploadClientFactory) TokenEnv: cfg.TokenEnv, Timeout: cfg.Timeout, newUploadClient: factory, + pollWait: waitForPoll, } } @@ -207,7 +209,11 @@ func (c *Client) Upload(ctx context.Context, req UploadRequest) (UploadResult, e Status: result.Status, UploadStatus: result.Status, } - status, statusErr := waitForRunStatus(runCtx, uploadClient, result.RunID, c.Timeout > 0) + pollWait := c.pollWait + if pollWait == nil { + pollWait = waitForPoll + } + status, statusErr := waitForRunStatus(runCtx, uploadClient, result.RunID, c.Timeout > 0, pollWait) status = sanitizeRunStatus(status) if status.RunID != "" || status.Status != "" { uploadResult.RunStatus = &RunStatus{ @@ -234,19 +240,15 @@ func (c *Client) Upload(ctx context.Context, req UploadRequest) (UploadResult, e return uploadResult, nil } -func waitForRunStatus(ctx context.Context, client uploadClient, runID string, poll bool) (runStatus, error) { +func waitForRunStatus(ctx context.Context, client uploadClient, runID string, poll bool, wait func(context.Context, time.Duration) error) (runStatus, error) { status, err := client.Status(ctx, runID) if err != nil || terminalRunStatus(status.Status) || !poll { return status, err } for { - timer := time.NewTimer(statusPollInterval) - select { - case <-ctx.Done(): - timer.Stop() - return status, fmt.Errorf("distributor run %q did not reach terminal status before timeout: %w", runID, ctx.Err()) - case <-timer.C: + if err := wait(ctx, statusPollInterval); err != nil { + return status, fmt.Errorf("distributor run %q did not reach terminal status before timeout: %w", runID, err) } next, err := client.Status(ctx, runID) @@ -260,6 +262,17 @@ func waitForRunStatus(ctx context.Context, client uploadClient, runID string, po } } +func waitForPoll(ctx context.Context, interval time.Duration) error { + timer := time.NewTimer(interval) + defer timer.Stop() + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + func terminalRunStatus(status string) bool { return status == "succeeded" || status == "failed" } diff --git a/internal/adapters/distributor/client_test.go b/internal/adapters/distributor/client_test.go index fba6dca..02e05e7 100644 --- a/internal/adapters/distributor/client_test.go +++ b/internal/adapters/distributor/client_test.go @@ -242,6 +242,7 @@ func TestUploadPollsUntilTerminalStatus(t *testing.T) { }, } client := newClient(cfg, factory.newClient) + client.pollWait = func(context.Context, time.Duration) error { return nil } result, err := client.Upload(context.Background(), validUploadRequest()) if err != nil { diff --git a/internal/briefing/spc_convective_outlooks_module_test.go b/internal/briefing/spc_convective_outlooks_module_test.go index 41e5c66..ca47b3e 100644 --- a/internal/briefing/spc_convective_outlooks_module_test.go +++ b/internal/briefing/spc_convective_outlooks_module_test.go @@ -221,15 +221,6 @@ func TestSPCOutlookBackgroundDefinitionAssetRecordsProvenance(t *testing.T) { } } -func TestSPCRiskDigestDefaultPolicyConstants(t *testing.T) { - if defaultSPCRiskDigestOutlookType != "categorical" { - t.Fatalf("defaultSPCRiskDigestOutlookType = %q, want categorical", defaultSPCRiskDigestOutlookType) - } - if defaultSPCRiskDigestMinimumSeverityRank != 3 { - t.Fatalf("defaultSPCRiskDigestMinimumSeverityRank = %d, want 3", defaultSPCRiskDigestMinimumSeverityRank) - } -} - func TestSPCConvectiveOutlooksRiskDigestFilters(t *testing.T) { tests := []struct { name string diff --git a/internal/collect/collect_test.go b/internal/collect/collect_test.go index 7d3ae97..9e00869 100644 --- a/internal/collect/collect_test.go +++ b/internal/collect/collect_test.go @@ -50,7 +50,7 @@ func TestRunWrapsAdapterConstructionError(t *testing.T) { } func TestRunWrapsFetchError(t *testing.T) { - server := collectionTestServer(t, map[string]int{"/observations": http.StatusBadGateway}) + server := collectionTestServer(t, map[string]int{"/observations": http.StatusBadRequest}) defer server.Close() cfg := config.Defaults()