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 b73042f2 fix(arrow/avro): stop reader workers before reuse (#1048)
b73042f2 is described below

commit b73042f2294d91bce92eafbd858968b80d02284f
Author: Minh Vu <[email protected]>
AuthorDate: Wed Aug 5 21:13:22 2026 +0200

    fix(arrow/avro): stop reader workers before reuse (#1048)
    
    ## Summary
    
    - make decoded datum and record batch sends cancellation-aware
    - wait for both reader workers before reuse or shutdown
    - drain and release queued record batches during close
    - synchronize reader errors and the decoded record counter
    - stop workers and release the record builder on final release
    
    ## Testing
    
    - `go test -race ./arrow/avro`
    
    Regression coverage fills one-slot decode and record queues without 
consuming them, then verifies that `Close` returns and all checked allocations 
are released. It also reuses a partially consumed reader and verifies that no 
records from the previous file enter the new stream.
---
 arrow/avro/loader.go      |  40 +++++++----
 arrow/avro/reader.go      | 110 +++++++++++++++++++++++++------
 arrow/avro/reader_test.go | 164 ++++++++++++++++++++++++++++++++++++++++++++++
 3 files changed, 282 insertions(+), 32 deletions(-)

diff --git a/arrow/avro/loader.go b/arrow/avro/loader.go
index fa97c426..5a3e4a73 100644
--- a/arrow/avro/loader.go
+++ b/arrow/avro/loader.go
@@ -20,6 +20,8 @@ import (
        "errors"
        "fmt"
        "io"
+
+       "github.com/apache/arrow-go/v18/arrow"
 )
 
 func (r *OCFReader) decodeOCFToChan() {
@@ -27,7 +29,7 @@ func (r *OCFReader) decodeOCFToChan() {
        for {
                select {
                case <-r.readerCtx.Done():
-                       r.err = fmt.Errorf("avro decoding cancelled, %d records 
read", r.avroDatumCount)
+                       r.setErr(fmt.Errorf("avro decoding cancelled, %d 
records read", r.avroDatumCount.Load()))
                        return
                default:
                        var datum any
@@ -36,17 +38,23 @@ func (r *OCFReader) decodeOCFToChan() {
                                if errors.Is(err, io.EOF) {
                                        return
                                }
-                               r.err = err
+                               r.setErr(err)
+                               return
+                       }
+                       select {
+                       case r.avroChan <- datum:
+                               r.avroDatumCount.Add(1)
+                       case <-r.readerCtx.Done():
+                               r.setErr(fmt.Errorf("avro decoding cancelled, 
%d records read", r.avroDatumCount.Load()))
                                return
                        }
-                       r.avroChan <- datum
-                       r.avroDatumCount++
                }
        }
 }
 
 func (r *OCFReader) recordFactory() {
        defer close(r.recChan)
+       defer close(r.bldDone)
        r.primed = true
        recChunk := 0
        switch {
@@ -54,12 +62,11 @@ func (r *OCFReader) recordFactory() {
                for data := range r.avroChan {
                        err := r.ldr.loadDatum(data)
                        if err != nil {
-                               r.err = err
+                               r.setErr(err)
                                return
                        }
                }
-               r.recChan <- r.bld.NewRecordBatch()
-               r.bldDone <- struct{}{}
+               r.sendRecord(r.bld.NewRecordBatch())
        case r.chunk >= 1:
                for data := range r.avroChan {
                        if recChunk == 0 {
@@ -67,18 +74,29 @@ func (r *OCFReader) recordFactory() {
                        }
                        err := r.ldr.loadDatum(data)
                        if err != nil {
-                               r.err = err
+                               r.setErr(err)
                                return
                        }
                        recChunk++
                        if recChunk >= r.chunk {
-                               r.recChan <- r.bld.NewRecordBatch()
+                               if !r.sendRecord(r.bld.NewRecordBatch()) {
+                                       return
+                               }
                                recChunk = 0
                        }
                }
                if recChunk != 0 {
-                       r.recChan <- r.bld.NewRecordBatch()
+                       r.sendRecord(r.bld.NewRecordBatch())
                }
-               r.bldDone <- struct{}{}
+       }
+}
+
+func (r *OCFReader) sendRecord(rec arrow.RecordBatch) bool {
+       select {
+       case r.recChan <- rec:
+               return true
+       case <-r.readerCtx.Done():
+               rec.Release()
+               return false
        }
 }
diff --git a/arrow/avro/reader.go b/arrow/avro/reader.go
index be1a926b..f3db1950 100644
--- a/arrow/avro/reader.go
+++ b/arrow/avro/reader.go
@@ -22,6 +22,7 @@ import (
        "errors"
        "fmt"
        "io"
+       "sync"
        "sync/atomic"
 
        "github.com/apache/arrow-go/v18/arrow"
@@ -60,6 +61,8 @@ type OCFReader struct {
        ldr    *dataLoader
        cur    arrow.RecordBatch
        err    error
+       errMu  sync.RWMutex
+       readWG sync.WaitGroup
 
        primed     bool
        readerCtx  context.Context
@@ -68,7 +71,7 @@ type OCFReader struct {
        maxRec     int
 
        avroChan       chan any
-       avroDatumCount int64
+       avroDatumCount atomic.Int64
        avroChanSize   int
        recChan        chan arrow.RecordBatch
 
@@ -127,16 +130,18 @@ func NewOCFReader(r io.Reader, opts ...Option) 
(*OCFReader, error) {
                rr.mem = memory.DefaultAllocator
        }
        rr.readerCtx, rr.readCancel = context.WithCancel(context.Background())
-       go rr.decodeOCFToChan()
-
-       rr.bld = array.NewRecordBuilder(rr.mem, rr.schema)
-       rr.bldMap = newFieldPos()
-       rr.ldr = newDataLoader()
-       for idx, fb := range rr.bld.Fields() {
-               mapFieldBuilders(fb, rr.schema.Field(idx), rr.bldMap)
-       }
-       rr.ldr.drawTree(rr.bldMap)
-       go rr.recordFactory()
+       rr.readWG.Add(1)
+       go func() {
+               defer rr.readWG.Done()
+               rr.decodeOCFToChan()
+       }()
+
+       rr.initBuilder()
+       rr.readWG.Add(1)
+       go func() {
+               defer rr.readWG.Done()
+               rr.recordFactory()
+       }()
        return rr, nil
 }
 
@@ -144,7 +149,7 @@ func NewOCFReader(r io.Reader, opts ...Option) (*OCFReader, 
error) {
 // new Avro file has an identical schema.
 func (rr *OCFReader) Reuse(r io.Reader, opts ...Option) error {
        rr.Close()
-       rr.err = nil
+       rr.clearErr()
        ocfr, err := ocf.NewReader(r)
        if err != nil {
                return fmt.Errorf("%w: could not create avro ocfreader", 
arrow.ErrInvalid)
@@ -169,10 +174,14 @@ func (rr *OCFReader) Reuse(r io.Reader, opts ...Option) 
error {
        for _, opt := range opts {
                opt(rr)
        }
+       if rr.bld != nil {
+               rr.bld.Release()
+       }
+       rr.initBuilder()
 
        rr.maxOCF = 0
        rr.maxRec = 0
-       rr.avroDatumCount = 0
+       rr.avroDatumCount.Store(0)
        rr.primed = false
 
        rr.avroChan = make(chan any, rr.avroChanSize)
@@ -180,14 +189,55 @@ func (rr *OCFReader) Reuse(r io.Reader, opts ...Option) 
error {
        rr.bldDone = make(chan struct{})
 
        rr.readerCtx, rr.readCancel = context.WithCancel(context.Background())
-       go rr.decodeOCFToChan()
-       go rr.recordFactory()
+       rr.readWG.Add(2)
+       go func() {
+               defer rr.readWG.Done()
+               rr.decodeOCFToChan()
+       }()
+       go func() {
+               defer rr.readWG.Done()
+               rr.recordFactory()
+       }()
        return nil
 }
 
+func (rr *OCFReader) initBuilder() {
+       if rr.mem == nil {
+               rr.mem = memory.DefaultAllocator
+       }
+       rr.bld = array.NewRecordBuilder(rr.mem, rr.schema)
+       rr.bldMap = newFieldPos()
+       rr.ldr = newDataLoader()
+       for idx, fb := range rr.bld.Fields() {
+               mapFieldBuilders(fb, rr.schema.Field(idx), rr.bldMap)
+       }
+       rr.ldr.drawTree(rr.bldMap)
+}
+
 // Err returns the last error encountered during the iteration over the
 // underlying Avro file.
-func (r *OCFReader) Err() error { return r.err }
+func (r *OCFReader) Err() error {
+       r.errMu.RLock()
+       defer r.errMu.RUnlock()
+       return r.err
+}
+
+func (r *OCFReader) setErr(err error) {
+       if err == nil {
+               return
+       }
+       r.errMu.Lock()
+       if r.err == nil {
+               r.err = err
+       }
+       r.errMu.Unlock()
+}
+
+func (r *OCFReader) clearErr() {
+       r.errMu.Lock()
+       r.err = nil
+       r.errMu.Unlock()
+}
 
 // AvroSchema returns the Avro schema of the Avro OCF
 func (r *OCFReader) AvroSchema() string { return r.avroSchema }
@@ -214,13 +264,29 @@ func (r *OCFReader) Metrics() string {
 }
 
 // OCFRecordsReadCount returns the number of Avro datum that were read from 
the Avro file.
-func (r *OCFReader) OCFRecordsReadCount() int64 { return r.avroDatumCount }
+func (r *OCFReader) OCFRecordsReadCount() int64 { return 
r.avroDatumCount.Load() }
 
 // Close closes the OCFReader's Avro record read cache and converted Arrow 
record cache. OCFReader must
 // be closed if the Avro OCF's records have not been read to completion.
 func (r *OCFReader) Close() {
+       if r.readCancel == nil {
+               return
+       }
        r.readCancel()
-       r.err = r.readerCtx.Err()
+       r.readWG.Wait()
+       if r.Err() == nil {
+               r.setErr(r.readerCtx.Err())
+       }
+       for rec := range r.recChan {
+               if rec != nil {
+                       rec.Release()
+               }
+       }
+       if r.cur != nil {
+               r.cur.Release()
+               r.cur = nil
+       }
+       r.readCancel = nil
 }
 
 func (r *OCFReader) editAvroSchema(e schemaEdit) error {
@@ -263,7 +329,7 @@ func (r *OCFReader) Next() bool {
                        r.cur = <-r.recChan
                }
        }
-       if r.err != nil {
+       if r.Err() != nil {
                return false
        }
 
@@ -342,8 +408,10 @@ func (r *OCFReader) Release() {
        debug.Assert(r.refs.Load() > 0, "too many releases")
 
        if r.refs.Add(-1) == 0 {
-               if r.cur != nil {
-                       r.cur.Release()
+               r.Close()
+               if r.bld != nil {
+                       r.bld.Release()
+                       r.bld = nil
                }
        }
 }
diff --git a/arrow/avro/reader_test.go b/arrow/avro/reader_test.go
index 726011ee..d3c25fdf 100644
--- a/arrow/avro/reader_test.go
+++ b/arrow/avro/reader_test.go
@@ -19,6 +19,7 @@ package avro
 import (
        "bytes"
        "encoding/json"
+       "errors"
        "fmt"
        "os"
        "path/filepath"
@@ -31,10 +32,22 @@ import (
        "github.com/apache/arrow-go/v18/arrow/extensions"
        "github.com/apache/arrow-go/v18/arrow/memory"
        "github.com/stretchr/testify/assert"
+       "github.com/stretchr/testify/require"
        "github.com/twmb/avro"
        "github.com/twmb/avro/ocf"
 )
 
+func TestOCFReaderPreservesFirstError(t *testing.T) {
+       reader := &OCFReader{}
+       first := errors.New("first error")
+
+       reader.setErr(first)
+       reader.setErr(errors.New("later error"))
+       reader.setErr(nil)
+
+       require.ErrorIs(t, reader.Err(), first)
+}
+
 func TestReader(t *testing.T) {
        tests := []struct {
                arrowSchema []arrow.Field
@@ -241,6 +254,157 @@ func TestReader(t *testing.T) {
 // A nullable logical timestamp must decode to a value rather than a null: the
 // union branch carries the logical type, so the reader has to honour it on the
 // branch and not just on a bare long.
+
+// TestOCFReaderBytesValues exercises avro `bytes` fields, both plain and as a
+// ["null","bytes"] union.
+func TestOCFReaderBytesValues(t *testing.T) {
+       schema := `{
+               "type": "record",
+               "name": "rec",
+               "fields": [
+                       {"name": "plain", "type": "bytes"},
+                       {"name": "nullable", "type": ["null", "bytes"]}
+               ]
+       }`
+       payload := []byte{0x00, 0x01, 0xfe, 0xff}
+
+       var buf bytes.Buffer
+       avroSchema, err := avro.Parse(schema)
+       assert.NoError(t, err)
+       enc, err := ocf.NewWriter(&buf, avroSchema)
+       assert.NoError(t, err)
+       assert.NoError(t, enc.Encode(map[string]any{
+               "plain":    payload,
+               "nullable": map[string]any{"bytes": payload},
+       }))
+       assert.NoError(t, enc.Encode(map[string]any{
+               "plain":    []byte{},
+               "nullable": nil,
+       }))
+       assert.NoError(t, enc.Close())
+
+       ar, err := NewOCFReader(bytes.NewReader(buf.Bytes()), WithChunk(-1))
+       assert.NoError(t, err)
+       defer ar.Close()
+
+       assert.True(t, ar.Next())
+       assert.NoError(t, ar.Err())
+       rec := ar.RecordBatch()
+
+       plain := rec.Column(0).(*array.Binary)
+       assert.Equal(t, payload, plain.Value(0))
+       assert.Equal(t, []byte{}, plain.Value(1))
+
+       nullable := rec.Column(1).(*array.Binary)
+       assert.Equal(t, payload, nullable.Value(0))
+       assert.True(t, nullable.IsNull(1))
+}
+
+func TestOCFReaderCloseUnblocksFullQueues(t *testing.T) {
+       const schema = 
`{"type":"record","name":"rec","fields":[{"name":"value","type":"long"}]}`
+       var buf bytes.Buffer
+       avroSchema, err := avro.Parse(schema)
+       assert.NoError(t, err)
+       enc, err := ocf.NewWriter(&buf, avroSchema)
+       assert.NoError(t, err)
+       for i := 0; i < 100; i++ {
+               assert.NoError(t, enc.Encode(map[string]any{"value": int64(i)}))
+       }
+       assert.NoError(t, enc.Close())
+
+       mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+       reader, err := NewOCFReader(bytes.NewReader(buf.Bytes()), 
WithAllocator(mem),
+               WithReadCacheSize(1), WithRecordCacheSize(1), WithChunk(1))
+       assert.NoError(t, err)
+
+       deadline := time.Now().Add(time.Second)
+       for reader.OCFRecordsReadCount() < 3 && time.Now().Before(deadline) {
+               time.Sleep(time.Millisecond)
+       }
+
+       done := make(chan struct{})
+       go func() {
+               reader.Close()
+               close(done)
+       }()
+       select {
+       case <-done:
+       case <-time.After(time.Second):
+               t.Fatal("Close blocked with full producer queues")
+       }
+
+       reader.Release()
+       mem.AssertSize(t, 0)
+}
+
+func TestOCFReaderReuseWaitsForPreviousWorkers(t *testing.T) {
+       const schema = 
`{"type":"record","name":"rec","fields":[{"name":"value","type":"long"}]}`
+       encode := func(start int64) []byte {
+               var buf bytes.Buffer
+               avroSchema, err := avro.Parse(schema)
+               assert.NoError(t, err)
+               enc, err := ocf.NewWriter(&buf, avroSchema)
+               assert.NoError(t, err)
+               for i := range int64(20) {
+                       assert.NoError(t, enc.Encode(map[string]any{"value": 
start + i}))
+               }
+               assert.NoError(t, enc.Close())
+               return buf.Bytes()
+       }
+
+       mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+       reader, err := NewOCFReader(bytes.NewReader(encode(0)), 
WithAllocator(mem),
+               WithReadCacheSize(1), WithRecordCacheSize(1), WithChunk(1))
+       assert.NoError(t, err)
+       assert.True(t, reader.Next())
+
+       assert.NoError(t, reader.Reuse(bytes.NewReader(encode(100))))
+       var values []int64
+       for reader.Next() {
+               values = append(values, 
reader.RecordBatch().Column(0).(*array.Int64).Value(0))
+       }
+       assert.NoError(t, reader.Err())
+       assert.Len(t, values, 20)
+       for i, value := range values {
+               assert.Equal(t, int64(100+i), value)
+       }
+
+       reader.Release()
+       mem.AssertSize(t, 0)
+}
+
+func TestOCFReaderReuseDiscardsPartialBuilderState(t *testing.T) {
+       const schema = 
`{"type":"record","name":"rec","fields":[{"name":"value","type":"long"}]}`
+       encode := func(value int64) []byte {
+               var buf bytes.Buffer
+               avroSchema, err := avro.Parse(schema)
+               require.NoError(t, err)
+               enc, err := ocf.NewWriter(&buf, avroSchema)
+               require.NoError(t, err)
+               require.NoError(t, enc.Encode(map[string]any{"value": value}))
+               require.NoError(t, enc.Close())
+               return buf.Bytes()
+       }
+
+       mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+       reader, err := NewOCFReader(bytes.NewReader(encode(0)), 
WithAllocator(mem))
+       require.NoError(t, err)
+       reader.Close()
+
+       reader.bld.Field(0).(*array.Int64Builder).Append(999)
+
+       require.NoError(t, reader.Reuse(bytes.NewReader(encode(100))))
+       require.True(t, reader.Next())
+       record := reader.RecordBatch()
+       require.EqualValues(t, 1, record.NumRows())
+       require.EqualValues(t, 100, record.Column(0).(*array.Int64).Value(0))
+       require.False(t, reader.Next())
+       require.NoError(t, reader.Err())
+
+       reader.Release()
+       mem.AssertSize(t, 0)
+}
+
 func TestOCFReaderNullableTimestamps(t *testing.T) {
        tests := []struct {
                logicalType string

Reply via email to