diff --git a/arrow/flight/flight_test.go b/arrow/flight/flight_test.go index 8d75aac2e..e6b8a25b9 100644 --- a/arrow/flight/flight_test.go +++ b/arrow/flight/flight_test.go @@ -24,6 +24,7 @@ import ( "sync" "sync/atomic" "testing" + "time" "github.com/apache/arrow-go/v18/arrow" "github.com/apache/arrow-go/v18/arrow/array" @@ -635,6 +636,42 @@ func TestStreamChunksFromReader_OK(t *testing.T) { } +type immediateErrorRecordReader struct { + err error + released atomic.Bool +} + +func (*immediateErrorRecordReader) Retain() {} +func (r *immediateErrorRecordReader) Release() { r.released.Store(true) } +func (*immediateErrorRecordReader) Schema() *arrow.Schema { return nil } +func (*immediateErrorRecordReader) Next() bool { return false } +func (*immediateErrorRecordReader) RecordBatch() arrow.RecordBatch { return nil } +func (*immediateErrorRecordReader) Record() arrow.RecordBatch { return nil } +func (r *immediateErrorRecordReader) Err() error { return r.err } + +func TestStreamChunksFromReader_CancellationWhileSendingError(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + rdr := &immediateErrorRecordReader{err: errors.New("read failed")} + ch := make(chan flight.StreamChunk) + done := make(chan struct{}) + go func() { + flight.StreamChunksFromReader(ctx, rdr, ch) + close(done) + }() + + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("StreamChunksFromReader blocked sending an error after cancellation") + } + + if !rdr.released.Load() { + t.Fatal("reader was not released") + } +} + // TestStreamChunksFromReader_HandlesCancellation verifies that context cancellation // causes StreamChunksFromReader to exit cleanly and release the reader. func TestStreamChunksFromReader_HandlesCancellation(t *testing.T) { diff --git a/arrow/flight/record_batch_reader.go b/arrow/flight/record_batch_reader.go index e6990a571..8c9d38765 100644 --- a/arrow/flight/record_batch_reader.go +++ b/arrow/flight/record_batch_reader.go @@ -248,8 +248,11 @@ func StreamChunksFromReader(ctx context.Context, rdr array.RecordReader, ch chan } if e, ok := rdr.(haserr); ok { - if e.Err() != nil { - ch <- StreamChunk{Err: e.Err()} + if err := e.Err(); err != nil { + select { + case ch <- StreamChunk{Err: err}: + case <-ctx.Done(): + } } } }