This is an automated email from the ASF dual-hosted git repository.
zeroshade pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-go.git
The following commit(s) were added to refs/heads/main by this push:
new 6fc85784 fix(arrow/ipc): make writer failures terminal (#1047)
6fc85784 is described below
commit 6fc85784bc86f42c8413af28340ed603e4d11a0a
Author: Minh Vu <[email protected]>
AuthorDate: Fri Aug 7 17:16:33 2026 +0200
fix(arrow/ipc): make writer failures terminal (#1047)
### Rationale for this change
An IPC writer can fail after writing part of a schema, dictionary, or
record payload. The writer must not continue with stale state or lose
the first failure.
### What changes are included in this PR?
Keep the first terminal error, mark the writer as started before schema
payloads are written, close a started payload writer exactly once,
preserve close and panic failures, and release retained dictionaries.
Closing without a schema still returns arrow.ErrInvalid.
### Are these changes tested?
- `go test ./arrow/ipc`
### Are there any user-facing changes?
Yes. After a terminal write or close failure, later writes return the
stored error instead of continuing with stale writer state.
---
arrow/ipc/writer.go | 54 ++++++++++++++---
arrow/ipc/writer_test.go | 154 +++++++++++++++++++++++++++++++++++++++++++++++
2 files changed, 199 insertions(+), 9 deletions(-)
diff --git a/arrow/ipc/writer.go b/arrow/ipc/writer.go
index a1c94877..b6c623a7 100644
--- a/arrow/ipc/writer.go
+++ b/arrow/ipc/writer.go
@@ -90,6 +90,7 @@ type Writer struct {
pw PayloadWriter
started bool
+ err error
schema *arrow.Schema
mapper dictutils.Mapper
codec flatbuf.CompressionType
@@ -136,10 +137,13 @@ func NewWriter(w io.Writer, opts ...Option) *Writer {
}
func (w *Writer) Close() error {
+ if w.err != nil {
+ return w.closeAfterFailure()
+ }
if !w.started {
err := w.start()
if err != nil {
- return err
+ return w.closeAfterFailure()
}
}
@@ -148,24 +152,47 @@ func (w *Writer) Close() error {
}
err := w.pw.Close()
+ w.pw = nil
+ w.releaseDictionaries()
if err != nil {
- return fmt.Errorf("arrow/ipc: could not close payload writer:
%w", err)
+ return w.fail(fmt.Errorf("arrow/ipc: could not close payload
writer: %w", err))
+ }
+
+ return nil
+}
+
+func (w *Writer) closeAfterFailure() error {
+ if w.started && w.pw != nil {
+ w.err = errors.Join(w.err, w.pw.Close())
}
+ w.releaseDictionaries()
w.pw = nil
+ return w.err
+}
+func (w *Writer) releaseDictionaries() {
for _, d := range w.lastWrittenDicts {
d.Release()
}
+ w.lastWrittenDicts = nil
+}
- return nil
+func (w *Writer) fail(err error) error {
+ if w.err == nil {
+ w.err = err
+ }
+ return w.err
}
func (w *Writer) Write(rec arrow.RecordBatch) (err error) {
defer func() {
if pErr := recover(); pErr != nil {
- err = utils.FormatRecoveredError("arrow/ipc: unknown
error while writing", pErr)
+ err = w.fail(utils.FormatRecoveredError("arrow/ipc:
unknown error while writing", pErr))
}
}()
+ if w.err != nil {
+ return w.err
+ }
incomingSchema := rec.Schema()
@@ -201,15 +228,18 @@ func (w *Writer) Write(rec arrow.RecordBatch) (err error)
{
err = writeDictionaryPayloads(w.mem, rec, false, w.emitDictDeltas,
&w.mapper, w.lastWrittenDicts, w.pw, enc)
if err != nil {
- return fmt.Errorf("arrow/ipc: failure writing dictionary
batches: %w", err)
+ return w.fail(fmt.Errorf("arrow/ipc: failure writing dictionary
batches: %w", err))
}
enc.reset()
if err := enc.Encode(&data, rec); err != nil {
- return fmt.Errorf("arrow/ipc: could not encode record to
payload: %w", err)
+ return w.fail(fmt.Errorf("arrow/ipc: could not encode record to
payload: %w", err))
}
- return w.pw.WritePayload(data)
+ if err := w.pw.WritePayload(data); err != nil {
+ return w.fail(err)
+ }
+ return nil
}
func writeDictionaryPayloads(mem memory.Allocator, batch arrow.RecordBatch,
isFileFormat bool, emitDictDeltas bool, mapper *dictutils.Mapper,
lastWrittenDicts map[int64]arrow.Array, pw PayloadWriter, encoder
*recordEncoder) error {
@@ -279,7 +309,12 @@ func writeDictionaryPayloads(mem memory.Allocator, batch
arrow.RecordBatch, isFi
}
func (w *Writer) start() error {
- w.started = true
+ if w.err != nil {
+ return w.err
+ }
+ if w.schema == nil {
+ return w.fail(fmt.Errorf("%w: cannot write IPC stream without a
schema", arrow.ErrInvalid))
+ }
w.mapper.ImportSchema(w.schema)
w.lastWrittenDicts = make(map[int64]arrow.Array)
@@ -288,10 +323,11 @@ func (w *Writer) start() error {
ps := payloadFromSchema(w.schema, w.mem, &w.mapper)
defer ps.Release()
+ w.started = true
for _, data := range ps {
err := w.pw.WritePayload(data)
if err != nil {
- return err
+ return w.fail(err)
}
}
diff --git a/arrow/ipc/writer_test.go b/arrow/ipc/writer_test.go
index 6de7ee0d..315787ad 100644
--- a/arrow/ipc/writer_test.go
+++ b/arrow/ipc/writer_test.go
@@ -19,6 +19,7 @@ package ipc
import (
"bytes"
"encoding/binary"
+ "errors"
"fmt"
"io"
"math"
@@ -35,6 +36,159 @@ import (
"github.com/apache/arrow-go/v18/arrow/memory"
)
+type failingPayloadWriter struct {
+ err error
+ closeErr error
+ failAfter int
+ payloads int
+ closeCall int
+}
+
+type shortWriteWriter struct{}
+
+func (shortWriteWriter) Write(p []byte) (int, error) {
+ return len(p) - 1, io.ErrShortWrite
+}
+
+func TestPayloadWriteRejectsShortWrites(t *testing.T) {
+ mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+ defer mem.AssertSize(t, 0)
+
+ bldr := array.NewRecordBuilder(mem,
arrow.NewSchema([]arrow.Field{{Name: "col", Type: arrow.PrimitiveTypes.Int8}},
nil))
+ bldr.Field(0).(*array.Int8Builder).Append(1)
+ rec := bldr.NewRecordBatch()
+ defer rec.Release()
+
+ payload, err := GetRecordBatchPayload(rec, WithAllocator(mem))
+ require.NoError(t, err)
+ defer payload.Release()
+
+ _, err = payload.WritePayload(shortWriteWriter{})
+ require.ErrorIs(t, err, io.ErrShortWrite)
+}
+
+func (w *failingPayloadWriter) Start() error { return nil }
+func (w *failingPayloadWriter) WritePayload(Payload) error {
+ w.payloads++
+ if w.failAfter == 0 || w.payloads >= w.failAfter {
+ return w.err
+ }
+ return nil
+}
+func (w *failingPayloadWriter) Close() error {
+ w.closeCall++
+ return w.closeErr
+}
+
+func TestWriterCloseFailureIsTerminal(t *testing.T) {
+ schema := arrow.NewSchema([]arrow.Field{{Name: "col", Type:
arrow.PrimitiveTypes.Int32}}, nil)
+ want := errors.New("close failed")
+ payloadWriter := &failingPayloadWriter{closeErr: want}
+ writer := NewWriterWithPayloadWriter(payloadWriter, WithSchema(schema))
+
+ require.ErrorIs(t, writer.Close(), want)
+ require.ErrorIs(t, writer.Close(), want)
+ require.Equal(t, 1, payloadWriter.closeCall)
+}
+
+func TestWriterSchemaFailureIsTerminal(t *testing.T) {
+ schema := arrow.NewSchema([]arrow.Field{{Name: "col", Type:
arrow.PrimitiveTypes.Int32}}, nil)
+ builder := array.NewRecordBuilder(memory.DefaultAllocator, schema)
+ defer builder.Release()
+ record := builder.NewRecordBatch()
+ defer record.Release()
+
+ want := errors.New("schema write failed")
+ payloadWriter := &failingPayloadWriter{err: want}
+ writer := NewWriterWithPayloadWriter(payloadWriter, WithSchema(schema))
+
+ require.ErrorIs(t, writer.Write(record), want)
+ require.ErrorIs(t, writer.Write(record), want)
+ require.Equal(t, 1, payloadWriter.payloads)
+ require.ErrorIs(t, writer.Close(), want)
+ require.Equal(t, 1, payloadWriter.closeCall)
+}
+
+func TestWriterCloseSchemaFailureClosesStartedPayloadWriter(t *testing.T) {
+ schema := arrow.NewSchema([]arrow.Field{{Name: "col", Type:
arrow.PrimitiveTypes.Int32}}, nil)
+ want := errors.New("schema write failed")
+ payloadWriter := &failingPayloadWriter{err: want}
+ writer := NewWriterWithPayloadWriter(payloadWriter, WithSchema(schema))
+
+ require.ErrorIs(t, writer.Close(), want)
+ require.Equal(t, 1, payloadWriter.payloads)
+ require.Equal(t, 1, payloadWriter.closeCall)
+}
+
+func TestWriterRecordEncodingFailureIsTerminal(t *testing.T) {
+ mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+ defer mem.AssertSize(t, 0)
+
+ deepType := arrow.PrimitiveTypes.Int32
+ for i := 0; i < kMaxNestingDepth+1; i++ {
+ deepType = arrow.ListOf(deepType)
+ }
+ jsonValue := strings.Repeat("[", kMaxNestingDepth+2) + "1" +
strings.Repeat("]", kMaxNestingDepth+2)
+ deepArray, _, err := array.FromJSON(mem, deepType,
strings.NewReader(jsonValue))
+ require.NoError(t, err)
+ defer deepArray.Release()
+
+ dictType := &arrow.DictionaryType{IndexType: arrow.PrimitiveTypes.Int8,
ValueType: arrow.BinaryTypes.String}
+ dictArray, _, err := array.FromJSON(mem, dictType,
strings.NewReader(`["value"]`))
+ require.NoError(t, err)
+ defer dictArray.Release()
+
+ schema := arrow.NewSchema([]arrow.Field{
+ {Name: "dict", Type: dictType},
+ {Name: "deep", Type: deepType},
+ }, nil)
+ record := array.NewRecordBatch(schema, []arrow.Array{dictArray,
deepArray}, 1)
+ defer record.Release()
+
+ payloadWriter := &failingPayloadWriter{}
+ writer := NewWriterWithPayloadWriter(payloadWriter, WithSchema(schema))
+
+ firstErr := writer.Write(record)
+ require.Error(t, firstErr)
+ require.Equal(t, 2, payloadWriter.payloads)
+
+ secondErr := writer.Write(record)
+ require.EqualError(t, secondErr, firstErr.Error())
+ require.Equal(t, 2, payloadWriter.payloads)
+ require.Error(t, writer.Close())
+}
+
+func TestWriterPayloadFailureClosesStartedPayloadWriter(t *testing.T) {
+ schema := arrow.NewSchema([]arrow.Field{{Name: "col", Type:
arrow.PrimitiveTypes.Int32}}, nil)
+ builder := array.NewRecordBuilder(memory.DefaultAllocator, schema)
+ defer builder.Release()
+ record := builder.NewRecordBatch()
+ defer record.Release()
+
+ payloadErr := errors.New("payload failed")
+ closeErr := errors.New("close failed")
+ payloadWriter := &failingPayloadWriter{err: payloadErr, closeErr:
closeErr, failAfter: 2}
+ writer := NewWriterWithPayloadWriter(payloadWriter, WithSchema(schema))
+
+ require.ErrorIs(t, writer.Write(record), payloadErr)
+ err := writer.Close()
+ require.ErrorIs(t, err, payloadErr)
+ require.ErrorIs(t, err, closeErr)
+ require.Equal(t, 2, payloadWriter.payloads)
+ require.Equal(t, 1, payloadWriter.closeCall)
+ require.ErrorIs(t, writer.Close(), payloadErr)
+ require.Equal(t, 1, payloadWriter.closeCall)
+}
+
+func TestWriterCloseWithoutSchemaReturnsError(t *testing.T) {
+ payloadWriter := &failingPayloadWriter{}
+ writer := NewWriterWithPayloadWriter(payloadWriter)
+
+ require.ErrorIs(t, writer.Close(), arrow.ErrInvalid)
+ require.Zero(t, payloadWriter.payloads)
+ require.Zero(t, payloadWriter.closeCall)
+}
+
// reproducer from ARROW-13529
func TestSliceAndWrite(t *testing.T) {
alloc := memory.NewGoAllocator()