diff --git a/CHANGELOG.md b/CHANGELOG.md index b6f097c..0bd249a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Changed + +- `workflows`: Change `Client.NewPollingTaskRunner` to accept a resolved `*Cluster`, executor, and logger directly; it no longer fetches a cluster or accepts task-runner options. +- `workflows`: Log recoverable task polling and result-reporting failures as warnings while they await retry. + +### Fixed + +- Retry recoverable HTTP transport failures, including EOFs, timeouts, and connection resets, with five total attempts and bounded exponential jitter. +- Allow polling task runners to finish graceful shutdown after reporting an active task's pending result. + ## [0.12.0] - 2026-09-01 ### Added diff --git a/internal/grpc/grpc_connect.go b/internal/grpc/grpc_connect.go index 1509944..f9bb813 100644 --- a/internal/grpc/grpc_connect.go +++ b/internal/grpc/grpc_connect.go @@ -4,33 +4,31 @@ import ( "bytes" "context" "encoding/json" - "errors" "io" "log/slog" + "math/rand/v2" "net/http" - "net/url" - "strings" "time" "connectrpc.com/connect" "github.com/hashicorp/go-retryablehttp" ) -// retryOnStatusUnavailable provides a retry policy for retrying requests if the server is unavailable. -func retryOnStatusUnavailable(ctx context.Context, resp *http.Response, err error) (bool, error) { +// shouldRetryTransientRequest retries recoverable transport failures, including EOFs, timeouts, and connection resets, +// as well as transient HTTP statuses. retryablehttp's standard policy excludes permanent transport configuration and +// certificate errors. +func shouldRetryTransientRequest(ctx context.Context, resp *http.Response, err error) (bool, error) { // do not retry on context.Canceled or context.DeadlineExceeded if ctx.Err() != nil { return false, ctx.Err() } if err != nil { - if v, ok := errors.AsType[*url.Error](err); ok { - // Retry if the error was due to a connection refused. - if strings.Contains(v.Error(), "connect: connection refused") { - slog.InfoContext(ctx, "Auth client retry", slog.Any("error", v)) - return true, v - } + shouldRetry, policyErr := retryablehttp.ErrorPropagatedRetryPolicy(ctx, resp, err) + if shouldRetry { + slog.InfoContext(ctx, "HTTP client retry", slog.Any("error", err)) } + return shouldRetry, policyErr } if resp != nil { @@ -52,7 +50,7 @@ func retryOnStatusUnavailable(ctx context.Context, resp *http.Response, err erro } if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 { - slog.InfoContext(ctx, "Auth client retry", + slog.InfoContext(ctx, "HTTP client retry", slog.String("status", resp.Status), slog.Int("status_code", resp.StatusCode), slog.String("protocol", resp.Proto), @@ -68,9 +66,23 @@ func RetryHTTPClient() connect.HTTPClient { retryClient.Logger = nil retryClient.RetryWaitMin = 20 * time.Millisecond retryClient.RetryWaitMax = 10 * time.Second - retryClient.RetryMax = 5 - retryClient.Backoff = retryablehttp.LinearJitterBackoff - retryClient.CheckRetry = retryOnStatusUnavailable + retryClient.RetryMax = 4 // Five total attempts: the initial request plus four retries. + retryClient.Backoff = exponentialJitterBackoff + retryClient.CheckRetry = shouldRetryTransientRequest return retryClient.StandardClient() } + +func exponentialJitterBackoff(minimum, maximum time.Duration, attempt int, _ *http.Response) time.Duration { + upperBound := minimum + for range attempt { + if upperBound >= maximum/2 { + upperBound = maximum + break + } + upperBound *= 2 + } + + lowerBound := upperBound / 2 + return lowerBound + rand.N(upperBound-lowerBound+1) +} diff --git a/internal/grpc/grpc_connect_test.go b/internal/grpc/grpc_connect_test.go new file mode 100644 index 0000000..eab5624 --- /dev/null +++ b/internal/grpc/grpc_connect_test.go @@ -0,0 +1,102 @@ +package grpc + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestShouldRetryTransientRequestRetriesTransportErrors(t *testing.T) { + tests := map[string]error{ + "EOF": io.EOF, + "unexpected EOF": io.ErrUnexpectedEOF, + "timeout": &url.Error{ + Op: "Post", + URL: "https://api.tilebox.dev", + Err: syscall.ETIMEDOUT, + }, + "connection refused": &url.Error{ + Op: "Post", + URL: "https://api.tilebox.dev", + Err: syscall.ECONNREFUSED, + }, + "connection reset": &url.Error{ + Op: "Post", + URL: "https://api.tilebox.dev", + Err: syscall.ECONNRESET, + }, + } + + for name, transportError := range tests { + t.Run(name, func(t *testing.T) { + shouldRetry, err := shouldRetryTransientRequest(context.Background(), nil, transportError) + + assert.True(t, shouldRetry) + require.NoError(t, err) + }) + } +} + +func TestShouldRetryTransientRequestStopsWhenContextIsCanceled(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + shouldRetry, err := shouldRetryTransientRequest(ctx, nil, io.EOF) + + assert.False(t, shouldRetry) + require.ErrorIs(t, err, context.Canceled) +} + +func TestRetryHTTPClientMakesFiveTotalAttempts(t *testing.T) { + var attempts atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) { + attempts.Add(1) + response.WriteHeader(http.StatusServiceUnavailable) + })) + t.Cleanup(server.Close) + + request, err := http.NewRequestWithContext(context.Background(), http.MethodGet, server.URL, nil) + require.NoError(t, err) + + response, err := RetryHTTPClient().Do(request) + if response != nil { + t.Cleanup(func() { require.NoError(t, response.Body.Close()) }) + } + + require.Error(t, err) + assert.Equal(t, int32(5), attempts.Load()) +} + +func TestExponentialJitterBackoffIsBounded(t *testing.T) { + const ( + minimum = 20 * time.Millisecond + maximum = 10 * time.Second + ) + tests := []struct { + attempt int + lowerBound time.Duration + upperBound time.Duration + }{ + {attempt: 0, lowerBound: 10 * time.Millisecond, upperBound: 20 * time.Millisecond}, + {attempt: 1, lowerBound: 20 * time.Millisecond, upperBound: 40 * time.Millisecond}, + {attempt: 2, lowerBound: 40 * time.Millisecond, upperBound: 80 * time.Millisecond}, + {attempt: 20, lowerBound: 5 * time.Second, upperBound: maximum}, + } + + for _, test := range tests { + for range 100 { + delay := exponentialJitterBackoff(minimum, maximum, test.attempt, nil) + assert.GreaterOrEqual(t, delay, test.lowerBound) + assert.LessOrEqual(t, delay, test.upperBound) + } + } +} diff --git a/workflows/v1/client.go b/workflows/v1/client.go index 8e796e0..178c54b 100644 --- a/workflows/v1/client.go +++ b/workflows/v1/client.go @@ -2,7 +2,7 @@ package workflows import ( "context" - "fmt" + "errors" "log/slog" "net" "net/http" @@ -80,22 +80,15 @@ func (c *Client) NewTaskRunner(ctx context.Context, options ...runner.Option) (* return newTaskRunner(ctx, c.taskService, c.Clusters, c.tracer, options...) } -// NewPollingTaskRunner creates a polling task runner for custom task executors. -func (c *Client) NewPollingTaskRunner(ctx context.Context, executor TaskExecutor, options ...runner.Option) (*PollingTaskRunner, error) { - opts := &runner.Options{ - ClusterSlug: "", - Logger: slog.Default(), - MeterProvider: otel.GetMeterProvider(), +// NewPollingTaskRunner creates a polling task runner for a resolved cluster and custom task executor. +func (c *Client) NewPollingTaskRunner(cluster *Cluster, executor TaskExecutor, logger *slog.Logger) (*PollingTaskRunner, error) { + if cluster == nil { + return nil, errors.New("cluster is required") } - for _, option := range options { - option(opts) - } - - cluster, err := c.Clusters.Get(ctx, opts.ClusterSlug) - if err != nil { - return nil, fmt.Errorf("failed to get cluster: %w", err) + if strings.TrimSpace(cluster.Slug) == "" { + return nil, errors.New("cluster slug is required") } - return NewPollingTaskRunner(c.taskService, cluster.Slug, executor, opts.Logger), nil + return NewPollingTaskRunner(c.taskService, cluster.Slug, executor, logger), nil } // clientConfig contains the configuration for Tilebox Workflows client. diff --git a/workflows/v1/client_polling_runner_test.go b/workflows/v1/client_polling_runner_test.go new file mode 100644 index 0000000..46d89ff --- /dev/null +++ b/workflows/v1/client_polling_runner_test.go @@ -0,0 +1,34 @@ +package workflows + +import ( + "log/slog" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestNewPollingTaskRunnerUsesResolvedCluster(t *testing.T) { + service := &mockMinimalTaskService{} + client := &Client{taskService: service} + executor := &failedResponseExecutor{} + logger := slog.New(slog.DiscardHandler) + cluster := &Cluster{Slug: "resolved-cluster"} + + pollingRunner, err := client.NewPollingTaskRunner(cluster, executor, logger) + + require.NoError(t, err) + require.Equal(t, "resolved-cluster", pollingRunner.clusterSlug) + require.Same(t, service, pollingRunner.service) + require.Same(t, executor, pollingRunner.executor) + require.Same(t, logger, pollingRunner.logger) +} + +func TestNewPollingTaskRunnerRequiresResolvedCluster(t *testing.T) { + client := &Client{} + + _, err := client.NewPollingTaskRunner(nil, &failedResponseExecutor{}, nil) + require.EqualError(t, err, "cluster is required") + + _, err = client.NewPollingTaskRunner(&Cluster{}, &failedResponseExecutor{}, nil) + require.EqualError(t, err, "cluster slug is required") +} diff --git a/workflows/v1/polling_runner.go b/workflows/v1/polling_runner.go index 15eaf98..f0109fd 100644 --- a/workflows/v1/polling_runner.go +++ b/workflows/v1/polling_runner.go @@ -167,6 +167,9 @@ func (r *PollingTaskRunner) run(ctx context.Context, stopWhenIdling bool) error // check whether the runner should still request a new task requestingNewTasks := r.requestNewTasks.Load() + if !requestingNewTasks && lastComputedTask == nil { + return nil + } // and get a list of all known task identifiers we are requesting work for taskIdentifiersCapableOfRunning := r.executor.TaskIdentifiers() @@ -213,6 +216,9 @@ func (r *PollingTaskRunner) run(ctx context.Context, stopWhenIdling bool) error switch outcome { case reportSucceeded: currentTaskLogContext = ctx + if !requestingNewTasks { + return nil + } case reportRetryNow: continue case reportRetryLater: @@ -236,7 +242,7 @@ func (r *PollingTaskRunner) run(ctx context.Context, stopWhenIdling bool) error taskResponse, err := r.service.NextTask(ctx, nil, nextTaskToRun) if err != nil { // easy retry case, since we don't have a task result to report anyway - logError(r.logger, ctx, err, "failed to request next task, will retry") + logFailure(r.logger, ctx, slog.LevelWarn, err, "failed to request next task, will retry") if r.idle(ctx, randomFallbackIdleDuration()) { return nil } @@ -250,7 +256,7 @@ func (r *PollingTaskRunner) run(ctx context.Context, stopWhenIdling bool) error task := work.GetNextTask() work = nil if isEmpty(task.GetId()) { - logError(r.logger, ctx, nil, "got a task without an ID - skipping to the next task") + logFailure(r.logger, ctx, slog.LevelError, nil, "got a task without an ID - skipping to the next task") continue } currentTaskLogContext = taskContextForLogs(ctx, task) @@ -392,7 +398,11 @@ func (r *PollingTaskRunner) reportPendingFailure(ctx, logCtx context.Context, pe taskFailedSimplified: taskFailedSimplified, } - logError(r.logger, logCtx, err, "failed to report task failure back to Tilebox", slog.Int("retry_count", retryAttempt), slog.Int("max_retries", maxTaskFailedRetries), slog.Bool("request_simplified", taskFailedSimplified), slog.Bool("resetting_runner", resetRunnerState)) + logLevel := slog.LevelWarn + if resetRunnerState { + logLevel = slog.LevelError + } + logFailure(r.logger, logCtx, logLevel, err, "failed to report task failure back to Tilebox", slog.Int("retry_count", retryAttempt), slog.Int("max_retries", maxTaskFailedRetries), slog.Bool("request_simplified", taskFailedSimplified), slog.Bool("resetting_runner", resetRunnerState)) if simplifiedRequest && !resetRunnerState { return pendingRetry, reportRetryNow } @@ -413,28 +423,28 @@ func (r *PollingTaskRunner) reportPendingComputed(ctx, logCtx context.Context, p return taskResponse, pendingReport{}, reportSucceeded } if shouldRetryRPCAfterTimeout(err) { - logError(r.logger, logCtx, err, "failed to report computed task, will retry", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) + logFailure(r.logger, logCtx, slog.LevelWarn, err, "failed to report computed task, will retry", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) return nil, pending, reportRetryLater } if nextTaskToRun != nil { - logError(r.logger, logCtx, err, "failed to report computed task and request next task due to request error, retrying without new work request", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) + logFailure(r.logger, logCtx, slog.LevelWarn, err, "failed to report computed task and request next task due to request error, retrying without new work request", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) taskResponse, err = r.service.NextTask(ctx, computedTask, nil) if err == nil { return taskResponse, pendingReport{}, reportSucceeded } if shouldRetryRPCAfterTimeout(err) || !shouldFailComputedTaskForRPC(err) { - logError(r.logger, logCtx, err, "failed to report computed task without requesting next task, will retry", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) + logFailure(r.logger, logCtx, slog.LevelWarn, err, "failed to report computed task without requesting next task, will retry", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) return nil, pending, reportRetryLater } } if !shouldFailComputedTaskForRPC(err) { - logError(r.logger, logCtx, err, "failed to report computed task due to request error, will retry", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) + logFailure(r.logger, logCtx, slog.LevelWarn, err, "failed to report computed task due to request error, will retry", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) return nil, pending, reportRetryLater } - logError(r.logger, logCtx, err, "failed to report computed task due to invalid payload, will report task as failed", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) + logFailure(r.logger, logCtx, slog.LevelError, err, "failed to report computed task due to invalid payload, will report task as failed", slog.String("task_id", computedTask.GetId().AsUUID().String()), slog.String("task_display", computedTask.GetDisplay())) return nil, pendingReport{ result: workflowsv1.ExecuteTaskResponse_builder{FailedTask: failedTaskFromComputedTaskRequestError(computedTask, err)}.Build(), @@ -484,7 +494,7 @@ func (r *PollingTaskRunner) extendTaskLease(ctx context.Context, taskID uuid.UUI r.logger.DebugContext(ctx, "extending task lease", slog.String("task_id", taskID.String()), slog.Duration("lease", lease), slog.Duration("wait", wait)) extension, err := r.service.ExtendTaskLease(ctx, taskID, 2*lease) if err != nil { - logError(r.logger, ctx, err, "failed to extend task lease", slog.String("task_id", taskID.String())) + logFailure(r.logger, ctx, slog.LevelError, err, "failed to extend task lease", slog.String("task_id", taskID.String())) return } if extension.GetLease() == nil { @@ -560,7 +570,7 @@ func (r *PollingTaskRunner) idle(ctx context.Context, duration time.Duration) bo } } -func logError(logger *slog.Logger, ctx context.Context, err error, msg string, args ...any) { +func logFailure(logger *slog.Logger, ctx context.Context, level slog.Level, err error, msg string, args ...any) { switch { case errors.Is(err, context.Canceled): return @@ -573,5 +583,5 @@ func logError(logger *slog.Logger, ctx context.Context, err error, msg string, a fields := make([]any, 0, len(args)+1) fields = append(fields, slog.Any("error", err)) fields = append(fields, args...) - logger.ErrorContext(ctx, msg, fields...) + logger.Log(ctx, level, msg, fields...) } diff --git a/workflows/v1/polling_runner_test.go b/workflows/v1/polling_runner_test.go index 2f2dc44..2b75ff6 100644 --- a/workflows/v1/polling_runner_test.go +++ b/workflows/v1/polling_runner_test.go @@ -1,9 +1,11 @@ package workflows import ( + "bytes" "context" "errors" "log/slog" + "strings" "testing" "time" @@ -135,6 +137,58 @@ func TestPollingTaskRunnerDoesNotFailComputedTaskWhenCombinedNextTaskRequestIsIn assert.Equal(t, taskID, service.computedTasks[0].GetId().AsUUID()) } +func TestPollingTaskRunnerStopsAfterReportingPendingComputedTask(t *testing.T) { + taskID := uuid.New() + taskDisplay := "Python task" + task := workflowsv1.Task_builder{ + Id: tileboxv1.NewUUID(taskID), + Identifier: workflowsv1.TaskIdentifier_builder{ + Name: "python.Task", + Version: "v1.0", + }.Build(), + State: workflowsv1.TaskState_TASK_STATE_RUNNING, + Display: &taskDisplay, + Lease: workflowsv1.TaskLease_builder{ + Lease: durationpb.New(5 * time.Minute), + RecommendedWaitUntilNextExtension: durationpb.New(5 * time.Minute), + }.Build(), + }.Build() + service := &requestErrorTaskService{nextTask: task} + executionStarted := make(chan struct{}) + finishExecution := make(chan struct{}) + executor := &blockingComputedResponseExecutor{ + executionStarted: executionStarted, + finishExecution: finishExecution, + response: workflowsv1.ExecuteTaskResponse_builder{ + ComputedTask: workflowsv1.ComputedTask_builder{Id: tileboxv1.NewUUID(taskID), Display: taskDisplay}.Build(), + }.Build(), + } + runner := NewPollingTaskRunner(service, "default", executor, slog.Default()) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan error, 1) + go func() { + done <- runner.RunForever(ctx) + }() + + <-executionStarted + runner.StopRequestingNewTasks() + close(finishExecution) + + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(time.Second): + cancel() + require.Fail(t, "runner did not stop after reporting its pending result") + } + require.Len(t, service.computedTasks, 1) + assert.Equal(t, taskID, service.computedTasks[0].GetId().AsUUID()) + require.Len(t, service.nextTaskRequests, 2) + assert.NotNil(t, service.nextTaskRequests[0]) + assert.Nil(t, service.nextTaskRequests[1]) +} + func TestPollingTaskRunnerRetriesTaskFailedRequestErrorOnceAsWorkflowFailure(t *testing.T) { taskID := uuid.New() taskDisplay := "Python task" @@ -226,6 +280,7 @@ func TestPollingTaskRunnerStopsRetryingTaskFailedAfterSecondRequestError(t *test func TestPollingTaskRunnerResetsTaskFailedAfterMaxRetryableErrors(t *testing.T) { taskID := uuid.New() + var logs bytes.Buffer service := &requestErrorTaskService{ taskFailedErrors: []error{ connect.NewError(connect.CodeUnavailable, errors.New("unavailable 1")), @@ -233,7 +288,7 @@ func TestPollingTaskRunnerResetsTaskFailedAfterMaxRetryableErrors(t *testing.T) connect.NewError(connect.CodeUnavailable, errors.New("unavailable 3")), }, } - runner := NewPollingTaskRunner(service, "default", &failedResponseExecutor{}, slog.Default()) + runner := NewPollingTaskRunner(service, "default", &failedResponseExecutor{}, slog.New(slog.NewTextHandler(&logs, nil))) pending := pendingReport{ result: workflowsv1.ExecuteTaskResponse_builder{ FailedTask: workflowsv1.TaskFailedRequest_builder{ @@ -251,6 +306,8 @@ func TestPollingTaskRunnerResetsTaskFailedAfterMaxRetryableErrors(t *testing.T) _, outcome = runner.reportPendingFailure(context.Background(), context.Background(), pending) assert.Equal(t, reportResetRunner, outcome) require.Len(t, service.taskFailedRequests, maxTaskFailedRetries) + assert.Equal(t, 2, strings.Count(logs.String(), "level=WARN")) + assert.Equal(t, 1, strings.Count(logs.String(), "level=ERROR")) } type failedResponseExecutor struct { @@ -267,8 +324,27 @@ func (e *failedResponseExecutor) ExecuteTask(context.Context, *workflowsv1.Task) return e.response, nil } +type blockingComputedResponseExecutor struct { + executionStarted chan<- struct{} + finishExecution <-chan struct{} + response *workflowsv1.ExecuteTaskResponse +} + +func (e *blockingComputedResponseExecutor) TaskIdentifiers() []*workflowsv1.TaskIdentifier { + return []*workflowsv1.TaskIdentifier{ + workflowsv1.TaskIdentifier_builder{Name: "python.Task", Version: "v1.0"}.Build(), + } +} + +func (e *blockingComputedResponseExecutor) ExecuteTask(context.Context, *workflowsv1.Task) (*workflowsv1.ExecuteTaskResponse, error) { + close(e.executionStarted) + <-e.finishExecution + return e.response, nil +} + type requestErrorTaskService struct { computedTasks []*workflowsv1.ComputedTask + nextTaskRequests []*workflowsv1.NextTaskToRun nextTask *workflowsv1.Task nextTaskComputedAndRequestErr error nextTaskComputedTaskErr error @@ -279,6 +355,7 @@ type requestErrorTaskService struct { var _ TaskService = &requestErrorTaskService{} func (s *requestErrorTaskService) NextTask(_ context.Context, computedTask *workflowsv1.ComputedTask, nextTaskToRun *workflowsv1.NextTaskToRun) (*workflowsv1.NextTaskResponse, error) { + s.nextTaskRequests = append(s.nextTaskRequests, nextTaskToRun) if computedTask != nil && nextTaskToRun != nil && s.nextTaskComputedAndRequestErr != nil { return nil, s.nextTaskComputedAndRequestErr }