diff --git a/writer/row_iterator.go b/writer/row_iterator.go index 5ea307d..ce62351 100644 --- a/writer/row_iterator.go +++ b/writer/row_iterator.go @@ -208,6 +208,10 @@ func WriteRowIterator(iter *spanner.RowIterator, w RowIteratorWriter) (*RowItera return RunRowIterator(iter, RowIteratorHooksFromWriter(w)) } +// errRowSourceDone is private so a yielded public iterator.Done cannot be +// mistaken for exhaustion. Facades return it only when their source ends. +var errRowSourceDone = errors.New("row source exhausted") + type rowIteratorFacade interface { next() (*spanner.Row, error) stop() @@ -220,7 +224,11 @@ type spannerRowIteratorFacade struct { } func (f spannerRowIteratorFacade) next() (*spanner.Row, error) { - return f.Next() + row, err := f.Next() + if err == iterator.Done { //nolint:errorlint // Wrapped or joined Done can carry a real failure. + return nil, errRowSourceDone + } + return row, err } func (f spannerRowIteratorFacade) stop() { @@ -269,7 +277,7 @@ func runRowIterator(fac rowIteratorFacade, hooks RowIteratorHooks) (*RowIterator first := true for { row, err := fac.next() - if err != nil && !errors.Is(err, iterator.Done) { + if err != nil && err != errRowSourceDone { //nolint:errorlint // Only the private, exact exhaustion sentinel ends successfully. return abort(err) } if first { @@ -280,7 +288,7 @@ func runRowIterator(fac rowIteratorFacade, hooks RowIteratorHooks) (*RowIterator } } } - if errors.Is(err, iterator.Done) { + if err == errRowSourceDone { //nolint:errorlint // Wrapped errors are failures, not exhaustion. break } if hooks.WriteRow != nil { diff --git a/writer/row_iterator_test.go b/writer/row_iterator_test.go index dde5112..b2aec54 100644 --- a/writer/row_iterator_test.go +++ b/writer/row_iterator_test.go @@ -8,7 +8,6 @@ import ( "cloud.google.com/go/spanner" sppb "cloud.google.com/go/spanner/apiv1/spannerpb" "github.com/google/go-cmp/cmp" - "google.golang.org/api/iterator" ) var ( @@ -36,7 +35,7 @@ func (s *stubRowIterator) next() (*spanner.Row, error) { s.i++ return row, nil } - return nil, iterator.Done + return nil, errRowSourceDone } func (s *stubRowIterator) stop() { diff --git a/writer/row_seq.go b/writer/row_seq.go index 0fbf574..de5d6cd 100644 --- a/writer/row_seq.go +++ b/writer/row_seq.go @@ -5,7 +5,6 @@ import ( "cloud.google.com/go/spanner" sppb "cloud.google.com/go/spanner/apiv1/spannerpb" - "google.golang.org/api/iterator" "github.com/apstndb/spanvalue" ) @@ -50,6 +49,8 @@ func RowSeq(rows ...*spanner.Row) iter.Seq2[*spanner.Row, error] { // A non-nil error yielded by rows aborts the run and is returned; the row // paired with it is ignored and the sequence is not consumed further. A nil // row yielded with a nil error aborts the run with [ErrNilRow]. +// Yielding [google.golang.org/api/iterator.Done], directly or wrapped, is also +// an error; only returning from the sequence indicates successful exhaustion. func RunRowSeq(md *sppb.ResultSetMetadata, rows iter.Seq2[*spanner.Row, error], hooks RowIteratorHooks) (*RowIteratorResult, error) { if rows == nil { return nil, ErrNilRowSeq @@ -126,7 +127,7 @@ type seqRowFacade struct { func (f *seqRowFacade) next() (*spanner.Row, error) { row, err, ok := f.nextPair() if !ok { - return nil, iterator.Done + return nil, errRowSourceDone } if err != nil { return nil, err diff --git a/writer/row_seq_source_error_test.go b/writer/row_seq_source_error_test.go new file mode 100644 index 0000000..61d4e53 --- /dev/null +++ b/writer/row_seq_source_error_test.go @@ -0,0 +1,62 @@ +package writer + +import ( + "errors" + "fmt" + "testing" + + "cloud.google.com/go/spanner" + sppb "cloud.google.com/go/spanner/apiv1/spannerpb" + "google.golang.org/api/iterator" +) + +func TestRunRowSeqPreservesDoneErrors(t *testing.T) { + t.Parallel() + for _, sourceErr := range []error{iterator.Done, fmt.Errorf("wrapped: %w", iterator.Done), errors.Join(errors.New("source failure"), iterator.Done)} { + for _, deferred := range []bool{false, true} { + for _, firstRow := range []bool{false, true} { + t.Run(fmt.Sprintf("%s/deferred=%v/firstRow=%v", sourceErr, deferred, firstRow), func(t *testing.T) { + t.Parallel() + row := mustNewSpannerRow(t, []string{"id"}, []any{int64(1)}) + md := metadataWithColumnNames("id") + var released, prepared, written, finished, afterError int + rows := func(yield func(*spanner.Row, error) bool) { + defer func() { released++ }() + if firstRow && !yield(row, nil) { + return + } + if !yield(row, sourceErr) { + return + } + afterError++ + } + hooks := RowIteratorHooks{ + PrepareMetadata: func(*sppb.ResultSetMetadata) error { prepared++; return nil }, + WriteRow: func(*spanner.Row) error { written++; return nil }, + Finish: func(*RowIteratorResult) error { finished++; return nil }, + } + var result *RowIteratorResult + var err error + if deferred { + result, err = RunRowSeqDeferredMetadata(func() *sppb.ResultSetMetadata { return md }, rows, hooks) + } else { + result, err = RunRowSeq(md, rows, hooks) + } + if err != sourceErr { //nolint:errorlint // Assert the original error identity, not just an errors.Is match. + t.Errorf("error identity changed: got %v, want %v", err, sourceErr) + } + wantRows := 0 + if firstRow { + wantRows = 1 + } + if released != 1 || finished != 0 || afterError != 0 || prepared != wantRows || written != wantRows { + t.Errorf("released=%d finished=%d afterError=%d prepared=%d written=%d", released, finished, afterError, prepared, written) + } + if result == nil || result.RowsRead != wantRows || result.Metadata != md { + t.Errorf("result=%+v", result) + } + }) + } + } + } +}