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
10 changes: 10 additions & 0 deletions pkg/database/database.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,16 @@ type FindResultInterface interface {
Err() error
}

// CursorResultInterface is the optional per-document iteration capability for
// find results. Consumers use it when individual decode failures must be
// handled without aborting the complete result set.
type CursorResultInterface interface {
FindResultInterface
Next(ctx context.Context) bool
Decode(dest any) error
Close(ctx context.Context) error
}

// UpdateResultInterface defines the interface for update operation results
type UpdateResultInterface interface {
MatchedCount() int64
Expand Down
48 changes: 46 additions & 2 deletions pkg/database/mock.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"encoding/json"
"fmt"
"reflect"
"time"

"go.mongodb.org/mongo-driver/bson"
Expand Down Expand Up @@ -86,10 +87,15 @@ type MockDeleteResult struct {

// MockFindResult implements FindResultInterface for testing
type MockFindResult struct {
results any
err error
results any
err error
index int
current any
currentSet bool
}

var _ CursorResultInterface = (*MockFindResult)(nil)

// All decodes all results into dest
func (m *MockFindResult) All(dest any) error {
if m.err != nil {
Expand All @@ -101,6 +107,44 @@ func (m *MockFindResult) All(dest any) error {
return copySliceResult(m.results, dest)
}

// Next advances to the next configured mock result.
func (m *MockFindResult) Next(context.Context) bool {
if m.err != nil || m.results == nil {
return false
}

results := reflect.ValueOf(m.results)
if results.Kind() != reflect.Slice && results.Kind() != reflect.Array {
return false
}
if m.index >= results.Len() {
m.current = nil
m.currentSet = false
return false
}

m.current = results.Index(m.index).Interface()
m.currentSet = true
m.index++
return true
}

// Decode decodes the current configured mock result.
func (m *MockFindResult) Decode(dest any) error {
if m.err != nil {
return m.err
}
if !m.currentSet {
return fmt.Errorf("find cursor has no current document")
}
return copyResult(m.current, dest)
}

// Close closes the mock cursor.
func (m *MockFindResult) Close(context.Context) error {
return m.err
}

// Err returns any error
func (m *MockFindResult) Err() error {
return m.err
Expand Down
56 changes: 56 additions & 0 deletions pkg/database/mock_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,62 @@ func TestMockDatabase(t *testing.T) {
}
})

t.Run("CursorIterationSkipsIndividualDecodeError", func(t *testing.T) {
mock := NewMockDatabase().QueueFind([]bson.M{
{"name": "first"},
{"name": []string{"invalid"}},
{"name": "third"},
}, nil)

result := mock.Find(context.Background(), "testdb", "users", bson.M{})
cursor, ok := result.(CursorResultInterface)
if !ok {
t.Fatal("mock find result must support cursor iteration")
}

var names []string
decodeErrors := 0
for cursor.Next(context.Background()) {
var document struct {
Name string `bson:"name"`
}
if err := cursor.Decode(&document); err != nil {
decodeErrors++
continue
}
names = append(names, document.Name)
}
if err := cursor.Err(); err != nil {
t.Fatalf("cursor error: %v", err)
}
if err := cursor.Close(context.Background()); err != nil {
t.Fatalf("close cursor: %v", err)
}
if decodeErrors != 1 || len(names) != 2 || names[0] != "first" || names[1] != "third" {
t.Fatalf("decode errors = %d, names = %v", decodeErrors, names)
}
})

t.Run("CursorIterationPreservesFindError", func(t *testing.T) {
expectedErr := errors.New("find failed")
result := NewMockDatabase().QueueFind(nil, expectedErr).
Find(context.Background(), "testdb", "users", bson.M{})
cursor := result.(CursorResultInterface)

if cursor.Next(context.Background()) {
t.Fatal("failed find must not have a next document")
}
if !errors.Is(cursor.Err(), expectedErr) {
t.Fatalf("cursor error = %v, want %v", cursor.Err(), expectedErr)
}
if !errors.Is(cursor.Decode(&bson.M{}), expectedErr) {
t.Fatalf("decode error = %v, want %v", cursor.Decode(&bson.M{}), expectedErr)
}
if !errors.Is(cursor.Close(context.Background()), expectedErr) {
t.Fatalf("close error = %v, want %v", cursor.Close(context.Background()), expectedErr)
}
})

t.Run("ExpectFindOneWithResult", func(t *testing.T) {
mock := NewMockDatabase()
expectedUser := map[string]any{
Expand Down
29 changes: 29 additions & 0 deletions pkg/database/mongodb.go
Original file line number Diff line number Diff line change
Expand Up @@ -382,6 +382,8 @@ type FindResult struct {
err error
}

var _ CursorResultInterface = (*FindResult)(nil)

// All decodes all results into the provided destination slice.
// The dest parameter must be a pointer to a slice.
func (fr *FindResult) All(dest any) error {
Expand All @@ -395,8 +397,35 @@ func (fr *FindResult) All(dest any) error {
return fr.cursor.All(fr.ctx, dest)
}

// Next advances to the next document in the result set.
func (fr *FindResult) Next(ctx context.Context) bool {
return fr.err == nil && fr.cursor != nil && fr.cursor.Next(ctx)
}

// Decode decodes the current document selected by Next.
func (fr *FindResult) Decode(dest any) error {
if fr.err != nil {
return fr.err
}
if fr.cursor == nil {
return fmt.Errorf("find cursor unavailable")
}
return fr.cursor.Decode(dest)
}

// Close closes the underlying MongoDB cursor.
func (fr *FindResult) Close(ctx context.Context) error {
if fr.cursor == nil {
return fr.err
}
return fr.cursor.Close(ctx)
}

// Err returns any error that occurred during the query.
func (fr *FindResult) Err() error {
if fr.err == nil && fr.cursor != nil {
return fr.cursor.Err()
}
return fr.err
}

Expand Down
Loading