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()

Reply via email to