Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 11 additions & 3 deletions writer/row_iterator.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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() {
Expand Down Expand Up @@ -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 {
Expand All @@ -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 {
Expand Down
3 changes: 1 addition & 2 deletions writer/row_iterator_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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() {
Expand Down
5 changes: 3 additions & 2 deletions writer/row_seq.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
62 changes: 62 additions & 0 deletions writer/row_seq_source_error_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}
}
}
Loading