Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
42 changes: 27 additions & 15 deletions internal/grpc/grpc_connect.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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),
Expand All @@ -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)
}
102 changes: 102 additions & 0 deletions internal/grpc/grpc_connect_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
}
23 changes: 8 additions & 15 deletions workflows/v1/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ package workflows

import (
"context"
"fmt"
"errors"
"log/slog"
"net"
"net/http"
Expand Down Expand Up @@ -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.
Expand Down
34 changes: 34 additions & 0 deletions workflows/v1/client_polling_runner_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
32 changes: 21 additions & 11 deletions workflows/v1/polling_runner.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down Expand Up @@ -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:
Expand All @@ -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
}
Expand All @@ -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)
Expand Down Expand Up @@ -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
}
Expand All @@ -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(),
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand All @@ -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...)
}
Loading