diff --git a/README.md b/README.md index e708836..6c1617a 100644 --- a/README.md +++ b/README.md @@ -88,6 +88,26 @@ opts.DataConverter = converter.NewCodecDataConverter(opts.DataConverter, lpc) temporalClient, _ := router.NewClient(opts) ``` +### Handling IO errors + +By default the codec panics on an IO error (a failed request, a bad status code, or a +checksum/size mismatch), so that a Temporal worker's `WorkflowPanicPolicy` can retry the +workflow task without risking non-determinism. See [this Temporal community +thread](https://community.temporal.io/t/panicing-within-a-dataconverter-and-or-payloadcodec/19305) +for why panicking is the safe default inside workflow code. + +Outside a Temporal worker — an API handler, a CLI, or a goroutine the Temporal SDK does not +manage — there is no such safety net, and an uncaught panic crashes the process. Pass +`WithoutPanicOnIOError()` when constructing the codec for one of those callers so it returns +a `*largepayloadcodec.IOError` instead: + +```golang +lpc, _ := largepayloadcodec.New( + largepayloadcodec.WithURL(lpsEndpoint), + largepayloadcodec.WithoutPanicOnIOError(), +) +``` + ## Architecture Architecturally, large payloads are passed through the `CodecDataConverter` which in turn uses the large payload codec to en- and decode the payloads. diff --git a/codec/codec.go b/codec/codec.go index c321097..e68822b 100644 --- a/codec/codec.go +++ b/codec/codec.go @@ -43,6 +43,10 @@ type Codec struct { skipUrlHealthCheck bool // disableEncoding when set to true encoding will be disabled. disableEncoding bool + // panicOnIOErrorDisabled when set to true, IO errors are returned as *IOError instead of panicking. + // Use this when the codec is used outside of a Temporal worker context (e.g. in an API handler or + // a goroutine not managed by the Temporal SDK), where panics are not caught and would crash the process. + panicOnIOErrorDisabled bool // customHeaders http headers to add the request sent to LargePayloadService customHeaders map[string][]string } @@ -168,6 +172,25 @@ func WithDecodeOnly() Option { }) } +// WithoutPanicOnIOError configures the codec to return an *IOError instead of panicking on IO errors. +// +// By default the codec panics on IO errors so that the Temporal worker's WorkflowPanicPolicy +// handles them, avoiding non-determinism issues inside workflow functions. +// +// Use this option when the codec is used outside of a Temporal worker context — for example +// in an API handler, a standalone client, or a goroutine the Temporal SDK does not manage (such +// as one spawned from an activity to call client.GetWorkflow(...).Get() in the background). +// Nothing catches a panic in those cases, so an IO error would crash the process. This is not +// needed for activity code itself: the SDK already recovers an activity's own panic into a +// retryable error, though returning this instead keeps that error's cause inspectable via +// errors.As. +func WithoutPanicOnIOError() Option { + return applier(func(c *Codec) error { + c.panicOnIOErrorDisabled = true + return nil + }) +} + // WithCustomHeader adds a custom header to append to the http headers of the request sent to LargePayloadService // when called with the same header it will not override the header but instead will append to its value. // fails when passed an empty header. @@ -259,6 +282,7 @@ func New(opts ...Option) (*Codec, error) { if err != nil { return nil, err } + defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("got status code %d from storage service at %s", resp.StatusCode, headURL) } @@ -324,21 +348,22 @@ func (c *Codec) encodePayload(ctx context.Context, payload *common.Payload) (*co addCustomHeaders(req, c.customHeaders) resp, err := c.client.Do(req) if err != nil { - panicOnIOError(err) // resp is nil when Do fails; panicOnIOError panics before resp.Body is accessed below + return nil, c.ioError(err) // resp is nil when Do fails; ioError panics/errors before resp.Body is accessed below } + defer func() { _ = resp.Body.Close() }() respBody, err := io.ReadAll(resp.Body) if err != nil { - panicOnIOError(err) + return nil, c.ioError(err) } if resp.StatusCode != http.StatusCreated && resp.StatusCode != http.StatusOK { - panicOnIOError(fmt.Errorf("server returned status code %d: %s", resp.StatusCode, respBody)) + return nil, c.ioError(fmt.Errorf("server returned status code %d: %s", resp.StatusCode, respBody)) } var key keyResponse if err := json.Unmarshal(respBody, &key); err != nil { - panicOnIOError(fmt.Errorf("unable to unmarshal put response: %w", err)) // key is zero-valued when unmarshal fails; panicOnIOError panics before key.Key is used below + return nil, c.ioError(fmt.Errorf("unable to unmarshal put response: %w", err)) // key is zero-valued when unmarshal fails; ioError panics/errors before key.Key is used below } result, err := converter.GetDefaultDataConverter().ToPayload(remotePayload{ @@ -410,27 +435,28 @@ func (c *Codec) decodePayload(ctx context.Context, payload *common.Payload, vers resp, err := c.client.Do(req) if err != nil { - panicOnIOError(err) // resp is nil when Do fails; panicOnIOError panics before resp.StatusCode is accessed below + return nil, c.ioError(err) // resp is nil when Do fails; ioError panics/errors before resp.StatusCode is accessed below } + defer func() { _ = resp.Body.Close() }() if resp.StatusCode != http.StatusOK { - panicOnIOError(fmt.Errorf("server returned status code %d", resp.StatusCode)) + return nil, c.ioError(fmt.Errorf("server returned status code %d", resp.StatusCode)) } sha2 := sha256.New() tee := io.TeeReader(resp.Body, sha2) b, err := io.ReadAll(tee) if err != nil { - panicOnIOError(err) + return nil, c.ioError(err) } if uint(len(b)) != remoteP.Size { - panicOnIOError(fmt.Errorf("wanted object of size %d, got %d", remoteP.Size, len(b))) + return nil, c.ioError(fmt.Errorf("wanted object of size %d, got %d", remoteP.Size, len(b))) } checkSum := hex.EncodeToString(sha2.Sum(nil)) if fmt.Sprintf("sha256:%s", checkSum) != remoteP.Digest { - panicOnIOError(fmt.Errorf("wanted object sha %s, got %s", remoteP.Digest, checkSum)) + return nil, c.ioError(fmt.Errorf("wanted object sha %s, got %s", remoteP.Digest, checkSum)) } return &common.Payload{ @@ -439,9 +465,28 @@ func (c *Codec) decodePayload(ctx context.Context, payload *common.Payload, vers }, nil } -// panicOnIOError panics the codec to force the workflows to handle the error via its [go.temporal.io/sdk/worker.WorkflowPanicPolicy]. -// If the codec returns an error, we can get into non-determinism issues. -// See https://community.temporal.io/t/panicing-within-a-dataconverter-and-or-payloadcodec/19305 -func panicOnIOError(err error) { - panic(fmt.Errorf("large payload codec IO error: %v", err)) +// IOError wraps a transport, HTTP-status, or integrity failure from the LPS server. +// Use errors.As to recover the cause. +type IOError struct { + Cause error +} + +func (e *IOError) Error() string { + return fmt.Sprintf("large payload codec IO error: %v", e.Cause) +} + +func (e *IOError) Unwrap() error { + return e.Cause +} + +// ioError panics by default so the Temporal worker's WorkflowPanicPolicy can handle it without +// workflow non-determinism (see +// https://community.temporal.io/t/panicing-within-a-dataconverter-and-or-payloadcodec/19305). +// WithoutPanicOnIOError returns the *IOError instead. +func (c *Codec) ioError(err error) error { + wrapped := &IOError{Cause: err} + if c.panicOnIOErrorDisabled { + return wrapped + } + panic(wrapped) } diff --git a/codec/codec_test.go b/codec/codec_test.go index b69b9ec..3a26bb2 100644 --- a/codec/codec_test.go +++ b/codec/codec_test.go @@ -7,10 +7,12 @@ package codec import ( "context" "errors" + "io" "net/http" "net/http/httptest" "os" "path/filepath" + "strings" "testing" "github.com/DataDog/temporal-large-payload-codec/server" @@ -22,6 +24,7 @@ import ( "github.com/golang/protobuf/proto" //nolint:staticcheck "github.com/stretchr/testify/require" "go.temporal.io/api/common/v1" + "go.temporal.io/sdk/converter" ) const ( @@ -404,14 +407,34 @@ func Test_io_error_in_put_panics(t *testing.T) { if r == nil { return } - _, ok := r.(error) - require.True(t, ok, "expected panic value to be an error") + _, ok := r.(*IOError) + require.True(t, ok, "expected panic value to be a *IOError") }() _, _ = c.Encode([]*common.Payload{&payload}) require.Fail(t, "expected panic but none occurred") } +func Test_io_error_in_put_returns_error_with_option(t *testing.T) { + d := NewPutRejectingDriver() + srv := httptest.NewServer(server.NewHttpHandler(d)) + defer srv.Close() + + c := setUpWithOptions(t, "v2", srv, WithoutPanicOnIOError()) + + payload := common.Payload{ + Metadata: map[string][]byte{ + "foo": []byte("bar"), + }, + Data: []byte("this is a longer message blah blah blah blah blah blah blah"), + } + + _, err := c.Encode([]*common.Payload{&payload}) + var ioErr *IOError + require.ErrorAs(t, err, &ioErr) + require.Contains(t, err.Error(), "large payload codec IO error") +} + func Test_io_error_in_get_panics(t *testing.T) { d := NewGetRejectingDriver() srv := httptest.NewServer(server.NewHttpHandler(d)) @@ -434,14 +457,45 @@ func Test_io_error_in_get_panics(t *testing.T) { if r == nil { return } - _, ok := r.(error) - require.True(t, ok, "expected panic value to be an error") + _, ok := r.(*IOError) + require.True(t, ok, "expected panic value to be a *IOError") }() _, _ = c.Decode(encodedPayloads) require.Fail(t, "expected panic but none occurred") } +func Test_io_error_in_get_returns_error_with_option(t *testing.T) { + d := &memory.Driver{} + srv := httptest.NewServer(server.NewHttpHandler(d)) + defer srv.Close() + + // Encode with a normal codec so the payload is stored, then decode with + // a WithoutPanicOnIOError codec backed by a driver that rejects gets. + encoder := setUpWithServer(t, "v2", srv, false) + + rejectingDriver := NewGetRejectingDriver() + rejectingDriver.memoryDriver = d + rejectingSrv := httptest.NewServer(server.NewHttpHandler(rejectingDriver)) + defer rejectingSrv.Close() + decoder := setUpWithOptions(t, "v2", rejectingSrv, WithoutPanicOnIOError()) + + payload := common.Payload{ + Metadata: map[string][]byte{ + "foo": []byte("bar"), + }, + Data: []byte("this is a longer message blah blah blah blah blah blah blah"), + } + + encodedPayloads, err := encoder.Encode([]*common.Payload{&payload}) + require.NoError(t, err) + + _, err = decoder.Decode(encodedPayloads) + var ioErr *IOError + require.ErrorAs(t, err, &ioErr) + require.Contains(t, err.Error(), "large payload codec IO error") +} + func setUp(t *testing.T, version string) (*httptest.Server, *Codec, storage.Driver) { d := &memory.Driver{} s := httptest.NewServer(server.NewHttpHandler(d)) @@ -450,6 +504,77 @@ func setUp(t *testing.T, version string) (*httptest.Server, *Codec, storage.Driv return s, c, d } +// closeTrackingBody wraps a response body to record whether Close was called. +type closeTrackingBody struct { + io.ReadCloser + closed *bool +} + +func (b *closeTrackingBody) Close() error { + *b.closed = true + return b.ReadCloser.Close() +} + +// rejectingTransport answers every request with a fixed error response, never making a real +// connection. Unlike an httptest.Server, nothing but the codec's own code can read or close +// the response body: a real net/http.Transport closes idle response bodies on its own in the +// background, which would let a test pass even if the codec itself never did. +type rejectingTransport struct { + closed *bool +} + +func (t *rejectingTransport) RoundTrip(_ *http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusInternalServerError, + Header: make(http.Header), + Body: &closeTrackingBody{ReadCloser: io.NopCloser(strings.NewReader("rejected")), closed: t.closed}, + }, nil +} + +func Test_io_error_closes_response_body(t *testing.T) { + payload := common.Payload{ + Metadata: map[string][]byte{"foo": []byte("bar")}, + Data: []byte("this is a longer message blah blah blah blah blah blah blah"), + } + + newRejectingCodec := func(t *testing.T, closed *bool) *Codec { + c, err := New( + WithURL("http://lps.invalid"), + WithNamespace("test"), + WithMinBytes(32), + WithoutUrlHealthCheck(), + WithoutPanicOnIOError(), + // An isolated client: WithHTTPRoundTripper mutates whatever client is + // already set, which defaults to the shared http.DefaultClient. + WithHTTPClient(&http.Client{}), + WithHTTPRoundTripper(&rejectingTransport{closed: closed}), + ) + require.NoError(t, err) + c.version = "v2" + return c + } + + t.Run("put", func(t *testing.T) { + closed := false + _, err := newRejectingCodec(t, &closed).Encode([]*common.Payload{&payload}) + require.Error(t, err) + require.True(t, closed, "expected the response body to be closed on an IO error") + }) + + t.Run("get", func(t *testing.T) { + closed := false + // A well-formed v2 reference is enough: decodePayload fails on the GET's status code, + // before it would ever need the referenced key to actually exist. + ref, err := converter.GetDefaultDataConverter().ToPayload(remotePayload{Key: "ignored", Digest: "sha256:ignored", Size: 1}) + require.NoError(t, err) + ref.Metadata[remoteCodecName] = []byte("v2") + + _, err = newRejectingCodec(t, &closed).Decode([]*common.Payload{ref}) + require.Error(t, err) + require.True(t, closed, "expected the response body to be closed on an IO error") + }) +} + func setUpWithServer(t *testing.T, version string, server *httptest.Server, withDecodeOnly bool) *Codec { opts := []Option{ WithURL(server.URL), @@ -468,6 +593,22 @@ func setUpWithServer(t *testing.T, version string, server *httptest.Server, with return c } +func setUpWithOptions(t *testing.T, version string, server *httptest.Server, extra ...Option) *Codec { + opts := []Option{ + WithURL(server.URL), + WithHTTPClient(server.Client()), + WithNamespace("test"), + WithMinBytes(32), + } + opts = append(opts, extra...) + c, err := New(opts...) + require.NoError(t, err) + + c.version = version + + return c +} + func fromFile(t *testing.T) []byte { path := filepath.Join("testdata", t.Name()) source, err := os.ReadFile(path)