This is an automated email from the ASF dual-hosted git repository.

laskoviymishka pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/iceberg-go.git


The following commit(s) were added to refs/heads/main by this push:
     new eb8285ce8 fix(table): cancel record writes on iterator stop (#1595)
eb8285ce8 is described below

commit eb8285ce8075a2f8784fbb3f118ec6567a005697
Author: Minh Vu <[email protected]>
AuthorDate: Thu Jul 30 00:12:23 2026 +0200

    fix(table): cancel record writes on iterator stop (#1595)
    
    ## What changed
    
    Cancel unpartitioned and fanout write contexts when a `WriteRecords`
    consumer stops early. Fanout writes now cancel record production
    independently from rolling writers and release retained batches left in
    the input queue after workers exit.
    
    Add unpartitioned and partitioned regressions showing that breaking
    after the first output stops the source before all records are consumed.
    The checked allocator also verifies retained Arrow memory is released.
    
    ## Why
    
    The iterator previously drained output without canceling production.
    Breaking iteration could therefore consume the entire source, continue
    writing files, or never return for an unbounded source.
    
    ## Testing
    
    - `go test ./table -run
    'Test(WriteRecords|FanoutWriter|PositionDeletePartitionedFanoutWriter)'
    -count=1 -timeout=180s`\n- `go test -race ./table -run
    'TestWriteRecords/TestEarlyStopCancelsRecordProduction' -count=1
    -timeout=120s`
    
    ---------
    
    Signed-off-by: Minh Vu <[email protected]>
---
 table/arrow_utils.go                               |  5 +-
 table/partitioned_fanout_writer.go                 | 30 +++++++--
 table/pos_delete_partitioned_fanout_writer.go      | 14 ++--
 table/pos_delete_partitioned_fanout_writer_test.go | 74 ++++++++++++++++++++++
 table/write_records_test.go                        | 58 +++++++++++++++++
 5 files changed, 170 insertions(+), 11 deletions(-)

diff --git a/table/arrow_utils.go b/table/arrow_utils.go
index 6e9e1e38c..86cbaa286 100644
--- a/table/arrow_utils.go
+++ b/table/arrow_utils.go
@@ -1866,12 +1866,13 @@ func recordsToDataFiles(ctx context.Context, 
rootLocation string, meta *Metadata
 func unpartitionedWrite(ctx context.Context, factory *writerFactory, records 
iter.Seq2[arrow.RecordBatch, error]) iter.Seq2[iceberg.DataFile, error] {
        outputCh := make(chan iceberg.DataFile, 1)
        errCh := make(chan error, 1)
+       writerCtx, cancel := context.WithCancel(ctx)
 
        go func() {
                defer close(outputCh)
                defer factory.stopCount()
 
-               writer := factory.newRollingDataWriter(ctx, "", nil, outputCh)
+               writer := factory.newRollingDataWriter(writerCtx, "", nil, 
outputCh)
                for rec, err := range records {
                        if err != nil {
                                errCh <- err
@@ -1880,6 +1881,7 @@ func unpartitionedWrite(ctx context.Context, factory 
*writerFactory, records ite
 
                                return
                        }
+
                        if err := writer.Add(rec); err != nil {
                                errCh <- err
                                close(errCh)
@@ -1899,6 +1901,7 @@ func unpartitionedWrite(ctx context.Context, factory 
*writerFactory, records ite
                        for range outputCh {
                        }
                }()
+               defer cancel()
                for df := range outputCh {
                        if !yield(df, nil) {
                                return
diff --git a/table/partitioned_fanout_writer.go 
b/table/partitioned_fanout_writer.go
index 3dda46732..39b716d02 100644
--- a/table/partitioned_fanout_writer.go
+++ b/table/partitioned_fanout_writer.go
@@ -80,8 +80,13 @@ func (p *partitionedFanoutWriter) Write(ctx context.Context, 
workers int) iter.S
        inputRecordsCh := make(chan arrow.RecordBatch, workers)
        outputDataFilesCh := make(chan iceberg.DataFile, workers)
 
-       fanoutWorkers, fanoutCtx := errgroup.WithContext(ctx)
+       fanoutBaseCtx, fanoutCancel := context.WithCancel(ctx)
+       fanoutWorkers, fanoutCtx := errgroup.WithContext(fanoutBaseCtx)
        writerCtx, writerCancel := context.WithCancel(ctx)
+       cancel := func() {
+               fanoutCancel()
+               writerCancel()
+       }
        startRecordFeeder(fanoutCtx, p.itr, fanoutWorkers, inputRecordsCh)
 
        for range workers {
@@ -90,7 +95,7 @@ func (p *partitionedFanoutWriter) Write(ctx context.Context, 
workers int) iter.S
                })
        }
 
-       return p.yieldDataFiles(fanoutWorkers, outputDataFilesCh, writerCancel)
+       return p.yieldDataFiles(fanoutWorkers, inputRecordsCh, 
outputDataFilesCh, cancel)
 }
 
 func startRecordFeeder(ctx context.Context, itr iter.Seq2[arrow.RecordBatch, 
error], fanoutWorkers *errgroup.Group, inputRecordsCh chan<- arrow.RecordBatch) 
{
@@ -175,32 +180,40 @@ func (p *partitionedFanoutWriter) processRecord(ctx 
context.Context, writerCtx c
        return nil
 }
 
-func (p *partitionedFanoutWriter) yieldDataFiles(fanoutWorkers 
*errgroup.Group, outputDataFilesCh chan iceberg.DataFile, writerCancel 
context.CancelFunc) iter.Seq2[iceberg.DataFile, error] {
+func (p *partitionedFanoutWriter) yieldDataFiles(fanoutWorkers 
*errgroup.Group, inputRecordsCh <-chan arrow.RecordBatch, outputDataFilesCh 
chan iceberg.DataFile, cancel context.CancelFunc) iter.Seq2[iceberg.DataFile, 
error] {
        return yieldDataFiles(
                p.writerFactory,
                fanoutWorkers,
+               inputRecordsCh,
                outputDataFilesCh,
                p.writerFactory.closeAll,
                p.writerFactory.abortAll,
-               writerCancel,
+               cancel,
        )
 }
 
 func yieldDataFiles(
        writerFactory *writerFactory,
        fanoutWorkers *errgroup.Group,
+       inputRecordsCh <-chan arrow.RecordBatch,
        outputDataFilesCh chan iceberg.DataFile,
        closeAll func() error,
        abortAll func(),
-       writerCancel context.CancelFunc,
+       cancel context.CancelFunc,
 ) iter.Seq2[iceberg.DataFile, error] {
        // Use a channel to safely communicate the error from the goroutine
        // to avoid a data race between writing err in the goroutine and 
reading it in the iterator.
        errCh := make(chan error, 1)
        go func() {
                defer close(outputDataFilesCh)
-               defer writerCancel()
+               defer cancel()
                err := fanoutWorkers.Wait()
+               // Wait includes the feeder, which closes inputRecordsCh, so 
draining cannot
+               // block. Any remaining batches were retained by the feeder but 
never dequeued;
+               // workers release dequeued batches themselves.
+               for record := range inputRecordsCh {
+                       record.Release()
+               }
                if err != nil {
                        abortAll()
                } else {
@@ -211,10 +224,15 @@ func yieldDataFiles(
        }()
 
        return func(yield func(iceberg.DataFile, error) bool) {
+               // LIFO defer order matters: cancel signals the producer first
+               // (synchronous, instant), then the drain pulls 
outputDataFilesCh so
+               // any in-flight stream send can complete and the producer's
+               // closeAll / fanoutWorkers.Wait paths unblock.
                defer func() {
                        for range outputDataFilesCh {
                        }
                }()
+               defer cancel()
 
                // Yield data files as they arrive - no error yet since 
goroutine is still running
                for f := range outputDataFilesCh {
diff --git a/table/pos_delete_partitioned_fanout_writer.go 
b/table/pos_delete_partitioned_fanout_writer.go
index f3bb24757..eb0dcad67 100644
--- a/table/pos_delete_partitioned_fanout_writer.go
+++ b/table/pos_delete_partitioned_fanout_writer.go
@@ -56,8 +56,13 @@ func (p *positionDeletePartitionedFanoutWriter) Write(ctx 
context.Context, worke
        inputRecordsCh := make(chan arrow.RecordBatch, workers)
        outputDataFilesCh := make(chan iceberg.DataFile, workers)
 
-       fanoutWorkers, fanoutCtx := errgroup.WithContext(ctx)
+       fanoutBaseCtx, fanoutCancel := context.WithCancel(ctx)
+       fanoutWorkers, fanoutCtx := errgroup.WithContext(fanoutBaseCtx)
        writerCtx, writerCancel := context.WithCancel(ctx)
+       cancel := func() {
+               fanoutCancel()
+               writerCancel()
+       }
        startRecordFeeder(fanoutCtx, p.itr, fanoutWorkers, inputRecordsCh)
 
        for range workers {
@@ -66,7 +71,7 @@ func (p *positionDeletePartitionedFanoutWriter) Write(ctx 
context.Context, worke
                })
        }
 
-       return p.yieldDataFiles(fanoutWorkers, outputDataFilesCh, writerCancel)
+       return p.yieldDataFiles(fanoutWorkers, inputRecordsCh, 
outputDataFilesCh, cancel)
 }
 
 func (p *positionDeletePartitionedFanoutWriter) fanout(ctx context.Context, 
writerCtx context.Context, inputRecordsCh <-chan arrow.RecordBatch, 
dataFilesChannel chan<- iceberg.DataFile) error {
@@ -143,13 +148,14 @@ func (p *positionDeletePartitionedFanoutWriter) 
partitionPath(partitionContext p
        return spec.PartitionToPath(data, schema), nil
 }
 
-func (p *positionDeletePartitionedFanoutWriter) yieldDataFiles(fanoutWorkers 
*errgroup.Group, outputDataFilesCh chan iceberg.DataFile, writerCancel 
context.CancelFunc) iter.Seq2[iceberg.DataFile, error] {
+func (p *positionDeletePartitionedFanoutWriter) yieldDataFiles(fanoutWorkers 
*errgroup.Group, inputRecordsCh <-chan arrow.RecordBatch, outputDataFilesCh 
chan iceberg.DataFile, cancel context.CancelFunc) iter.Seq2[iceberg.DataFile, 
error] {
        return yieldDataFiles(
                p.writerFactory,
                fanoutWorkers,
+               inputRecordsCh,
                outputDataFilesCh,
                p.writerFactory.closeAll,
                p.writerFactory.abortAll,
-               writerCancel,
+               cancel,
        )
 }
diff --git a/table/pos_delete_partitioned_fanout_writer_test.go 
b/table/pos_delete_partitioned_fanout_writer_test.go
index cace2762f..79ffe7e84 100644
--- a/table/pos_delete_partitioned_fanout_writer_test.go
+++ b/table/pos_delete_partitioned_fanout_writer_test.go
@@ -25,10 +25,14 @@ import (
        "runtime"
        "slices"
        "strings"
+       "sync/atomic"
        "testing"
        "time"
 
        "github.com/apache/arrow-go/v18/arrow"
+       "github.com/apache/arrow-go/v18/arrow/array"
+       "github.com/apache/arrow-go/v18/arrow/compute"
+       "github.com/apache/arrow-go/v18/arrow/memory"
        "github.com/apache/iceberg-go"
        "github.com/apache/iceberg-go/internal"
        "github.com/apache/iceberg-go/io"
@@ -402,6 +406,76 @@ func 
TestPositionDeletePartitionedFanoutWriterRoutesPartitionsIndependently(t *t
        assert.Equal(t, int64(1), byPart[2].Count(), "id=2 delete file must 
contain only the one row targeting pathB")
 }
 
+func 
TestPositionDeletePartitionedFanoutWriterEarlyStopCancelsRecordProduction(t 
*testing.T) {
+       t.Parallel()
+
+       const path = "file://t/id=1/a.parquet"
+       mem := memory.NewCheckedAllocator(memory.DefaultAllocator)
+       defer mem.AssertSize(t, 0)
+       ctx := compute.WithAllocator(t.Context(), mem)
+
+       tableSchema := iceberg.NewSchema(
+               0,
+               iceberg.NestedField{ID: 1, Name: "id", Type: 
iceberg.PrimitiveTypes.Int32, Required: true},
+       )
+       partitionSpec := iceberg.NewPartitionSpec(iceberg.PartitionField{
+               FieldID: 1000, SourceIDs: []int{1}, Name: "id", Transform: 
iceberg.IdentityTransform{},
+       })
+       metadataBuilder, err := NewMetadataBuilder(2)
+       require.NoError(t, err)
+       require.NoError(t, metadataBuilder.AddSchema(tableSchema))
+       require.NoError(t, metadataBuilder.SetCurrentSchemaID(0))
+       require.NoError(t, metadataBuilder.AddPartitionSpec(&partitionSpec, 
true))
+       require.NoError(t, metadataBuilder.SetDefaultSpecID(0))
+       require.NoError(t, metadataBuilder.AddSortOrder(&UnsortedSortOrder))
+       require.NoError(t, metadataBuilder.SetDefaultSortOrderID(0))
+       latestMeta, err := metadataBuilder.Build()
+       require.NoError(t, err)
+
+       var produced atomic.Int32
+       records := func(yield func(arrow.RecordBatch, error) bool) {
+               for range 1000 {
+                       produced.Add(1)
+                       batch, _, err := array.RecordFromJSON(
+                               mem,
+                               PositionalDeleteArrowSchema,
+                               
strings.NewReader(fmt.Sprintf(`[{"file_path":%q,"pos":0}]`, path)),
+                       )
+                       require.NoError(t, err)
+                       accepted := yield(batch, nil)
+                       batch.Release()
+                       if !accepted {
+                               return
+                       }
+               }
+       }
+
+       writeUUID := uuid.New()
+       factory, err := newWriterFactory(t.TempDir(), recordWritingArgs{
+               fs: &io.LocalFS{}, sc: PositionalDeleteArrowSchema, writeUUID: 
&writeUUID, counter: internal.Counter(0),
+       }, metadataBuilder, iceberg.PositionalDeleteSchema, 1,
+               withContentType(iceberg.EntryContentPosDeletes),
+               withFactoryFileSchema(iceberg.PositionalDeleteSchema))
+       require.NoError(t, err)
+       writer := newPositionDeletePartitionedFanoutWriter(
+               latestMeta,
+               map[string]partitionContext{path: {partitionData: 
map[int]any{1000: int32(1)}, specID: 0}},
+               records,
+               factory,
+       )
+
+       for dataFile, writeErr := range writer.Write(ctx, 1) {
+               require.NoError(t, writeErr)
+               require.NotNil(t, dataFile)
+
+               break
+       }
+
+       assert.Positive(t, produced.Load())
+       assert.Less(t, produced.Load(), int32(100))
+       require.Zero(t, mem.CurrentAlloc())
+}
+
 func TestPositionDeletePartitionedNoGoroutineLeak(t *testing.T) {
        t.Parallel()
 
diff --git a/table/write_records_test.go b/table/write_records_test.go
index bfe5048e2..a2d2ab0f8 100644
--- a/table/write_records_test.go
+++ b/table/write_records_test.go
@@ -23,6 +23,7 @@ import (
        "iter"
        "path/filepath"
        "strings"
+       "sync/atomic"
        "testing"
 
        "github.com/apache/arrow-go/v18/arrow"
@@ -233,6 +234,63 @@ func (s *WriteRecordsTestSuite) 
TestSmallTargetFileSizeProducesMultipleFiles() {
        s.Equal(int64(1000), totalRows)
 }
 
+func (s *WriteRecordsTestSuite) TestEarlyStopCancelsRecordProduction() {
+       tests := []struct {
+               name        string
+               partitioned bool
+       }{
+               {name: "unpartitioned"},
+               {name: "partitioned", partitioned: true},
+       }
+
+       for _, tt := range tests {
+               s.Run(tt.name, func() {
+                       loc := filepath.ToSlash(s.T().TempDir())
+                       iceSch := iceberg.NewSchema(1,
+                               iceberg.NestedField{ID: 1, Name: "id", Type: 
iceberg.PrimitiveTypes.Int32},
+                               iceberg.NestedField{ID: 2, Name: "name", Type: 
iceberg.PrimitiveTypes.String},
+                       )
+                       spec := iceberg.NewPartitionSpec()
+                       if tt.partitioned {
+                               spec = 
iceberg.NewPartitionSpec(iceberg.PartitionField{
+                                       SourceIDs: []int{1}, FieldID: 1000, 
Name: "id", Transform: iceberg.IdentityTransform{},
+                               })
+                       }
+                       meta, err := table.NewMetadata(iceSch, &spec, 
table.UnsortedSortOrder, loc, iceberg.Properties{})
+                       s.Require().NoError(err)
+                       tbl := table.New(
+                               table.Identifier{"test", tt.name}, meta, 
filepath.Join(loc, "metadata", "v1.metadata.json"),
+                               func(context.Context) (iceio.IO, error) { 
return iceio.LocalFS{}, nil }, nil,
+                       )
+
+                       schema := s.arrowSchema()
+                       var produced atomic.Int32
+                       records := func(yield func(arrow.RecordBatch, error) 
bool) {
+                               for range 1000 {
+                                       produced.Add(1)
+                                       record := s.buildRecords(schema, 100)
+                                       accepted := yield(record, nil)
+                                       if !accepted {
+                                               return
+                                       }
+                               }
+                       }
+
+                       for df, writeErr := range table.WriteRecords(
+                               s.ctx, tbl, schema, records, 
table.WithTargetFileSize(1), table.WithMaxWriteWorkers(1),
+                       ) {
+                               s.Require().NoError(writeErr)
+                               s.Require().NotNil(df)
+
+                               break
+                       }
+
+                       s.Positive(produced.Load())
+                       s.Less(produced.Load(), int32(100))
+               })
+       }
+}
+
 func (s *WriteRecordsTestSuite) TestWithWriteUUID() {
        loc := filepath.ToSlash(s.T().TempDir())
        tbl := s.newTable(loc)

Reply via email to