From 271bfa288b3df86b37e5857ff26046cf50fd6c3f Mon Sep 17 00:00:00 2001 From: Kilian Boute Date: Thu, 13 Aug 2026 14:09:32 +0000 Subject: [PATCH] fix(database): harden insert contracts Forward MongoDB insert options and add configurable insert responses, queues, call tracking, reset support, and tests to MockDatabase. --- pkg/database/mock.go | 102 +++++++++++++++++++++++++++++++++++ pkg/database/mock_test.go | 85 +++++++++++++++++++++++++++++ pkg/database/mongodb.go | 34 ++++++++++-- pkg/database/mongodb_test.go | 23 ++++++++ 4 files changed, 240 insertions(+), 4 deletions(-) diff --git a/pkg/database/mock.go b/pkg/database/mock.go index 4c419d7..aea4d8d 100644 --- a/pkg/database/mock.go +++ b/pkg/database/mock.go @@ -40,6 +40,12 @@ type MockDatabase struct { // CountFunc allows customizing Count behavior CountFunc func(ctx context.Context, db string, collection string, filter any, opts ...any) (int64, error) + // InsertOneFunc allows customizing InsertOne behavior + InsertOneFunc func(ctx context.Context, db string, collection string, document any, opts ...any) (any, error) + + // InsertManyFunc allows customizing InsertMany behavior + InsertManyFunc func(ctx context.Context, db string, collection string, documents []any, opts ...any) (any, error) + // Sequential response queues for multiple calls PingQueue []PingResponse FindQueue []FindResponse @@ -50,6 +56,8 @@ type MockDatabase struct { DeleteOneQueue []DeleteResponse DeleteManyQueue []DeleteResponse CountQueue []CountResponse + InsertOneQueue []InsertResponse + InsertManyQueue []InsertResponse // Call tracking PingCalls []PingCall @@ -61,6 +69,8 @@ type MockDatabase struct { DeleteOneCalls []DeleteCall DeleteManyCalls []DeleteCall CountCalls []CountCall + InsertOneCalls []InsertOneCall + InsertManyCalls []InsertManyCall } var _ DatabaseInterface = (*MockDatabase)(nil) @@ -301,6 +311,12 @@ type CountResponse struct { Err error } +// InsertResponse represents a queued insert response. +type InsertResponse struct { + Result any + Err error +} + // CountCall records a call to Count type CountCall struct { Ctx context.Context @@ -310,6 +326,24 @@ type CountCall struct { Opts []any } +// InsertOneCall records a call to InsertOne. +type InsertOneCall struct { + Ctx context.Context + Db string + Collection string + Document any + Opts []any +} + +// InsertManyCall records a call to InsertMany. +type InsertManyCall struct { + Ctx context.Context + Db string + Collection string + Documents []any + Opts []any +} + // NewMockDatabase creates a new MockDatabase with sensible defaults func NewMockDatabase() *MockDatabase { return &MockDatabase{ @@ -349,6 +383,8 @@ func NewMockDatabase() *MockDatabase { DeleteOneCalls: []DeleteCall{}, DeleteManyCalls: []DeleteCall{}, CountCalls: []CountCall{}, + InsertOneCalls: []InsertOneCall{}, + InsertManyCalls: []InsertManyCall{}, PingQueue: []PingResponse{}, FindQueue: []FindResponse{}, FindOneQueue: []FindOneResponse{}, @@ -358,6 +394,8 @@ func NewMockDatabase() *MockDatabase { DeleteOneQueue: []DeleteResponse{}, DeleteManyQueue: []DeleteResponse{}, CountQueue: []CountResponse{}, + InsertOneQueue: []InsertResponse{}, + InsertManyQueue: []InsertResponse{}, } } @@ -713,11 +751,43 @@ func (m *MockDatabase) Count(ctx context.Context, db string, collection string, // InsertOne implements DatabaseInterface func (m *MockDatabase) InsertOne(ctx context.Context, db string, collection string, document any, opts ...any) (any, error) { + m.InsertOneCalls = append(m.InsertOneCalls, InsertOneCall{ + Ctx: ctx, + Db: db, + Collection: collection, + Document: document, + Opts: opts, + }) + + if len(m.InsertOneQueue) > 0 { + response := m.InsertOneQueue[0] + m.InsertOneQueue = m.InsertOneQueue[1:] + return response.Result, response.Err + } + if m.InsertOneFunc != nil { + return m.InsertOneFunc(ctx, db, collection, document, opts...) + } return nil, fmt.Errorf("InsertOne not implemented in MockDatabase") } // InsertMany implements DatabaseInterface func (m *MockDatabase) InsertMany(ctx context.Context, db string, collection string, documents []any, opts ...any) (any, error) { + m.InsertManyCalls = append(m.InsertManyCalls, InsertManyCall{ + Ctx: ctx, + Db: db, + Collection: collection, + Documents: documents, + Opts: opts, + }) + + if len(m.InsertManyQueue) > 0 { + response := m.InsertManyQueue[0] + m.InsertManyQueue = m.InsertManyQueue[1:] + return response.Result, response.Err + } + if m.InsertManyFunc != nil { + return m.InsertManyFunc(ctx, db, collection, documents, opts...) + } return nil, fmt.Errorf("InsertMany not implemented in MockDatabase") } @@ -732,6 +802,8 @@ func (m *MockDatabase) Reset() { m.DeleteOneCalls = []DeleteCall{} m.DeleteManyCalls = []DeleteCall{} m.CountCalls = []CountCall{} + m.InsertOneCalls = []InsertOneCall{} + m.InsertManyCalls = []InsertManyCall{} m.PingQueue = []PingResponse{} m.FindQueue = []FindResponse{} m.FindOneQueue = []FindOneResponse{} @@ -741,6 +813,8 @@ func (m *MockDatabase) Reset() { m.DeleteOneQueue = []DeleteResponse{} m.DeleteManyQueue = []DeleteResponse{} m.CountQueue = []CountResponse{} + m.InsertOneQueue = []InsertResponse{} + m.InsertManyQueue = []InsertResponse{} } // ExpectPing sets up an expectation for Ping @@ -883,8 +957,36 @@ func (m *MockDatabase) ExpectCount(count int64, err error) *MockDatabase { return m } +// ExpectInsertOne sets up an InsertOne response. +func (m *MockDatabase) ExpectInsertOne(result any, err error) *MockDatabase { + m.InsertOneFunc = func(ctx context.Context, db string, collection string, document any, opts ...any) (any, error) { + return result, err + } + return m +} + +// ExpectInsertMany sets up an InsertMany response. +func (m *MockDatabase) ExpectInsertMany(result any, err error) *MockDatabase { + m.InsertManyFunc = func(ctx context.Context, db string, collection string, documents []any, opts ...any) (any, error) { + return result, err + } + return m +} + // QueueCount adds a Count response to the queue for sequential calls func (m *MockDatabase) QueueCount(count int64, err error) *MockDatabase { m.CountQueue = append(m.CountQueue, CountResponse{Count: count, Err: err}) return m } + +// QueueInsertOne adds an InsertOne response to the queue. +func (m *MockDatabase) QueueInsertOne(result any, err error) *MockDatabase { + m.InsertOneQueue = append(m.InsertOneQueue, InsertResponse{Result: result, Err: err}) + return m +} + +// QueueInsertMany adds an InsertMany response to the queue. +func (m *MockDatabase) QueueInsertMany(result any, err error) *MockDatabase { + m.InsertManyQueue = append(m.InsertManyQueue, InsertResponse{Result: result, Err: err}) + return m +} diff --git a/pkg/database/mock_test.go b/pkg/database/mock_test.go index 39f1070..5633931 100644 --- a/pkg/database/mock_test.go +++ b/pkg/database/mock_test.go @@ -826,3 +826,88 @@ func TestMockDatabaseCount(t *testing.T) { } }) } + +func TestMockDatabaseInsert(t *testing.T) { + t.Run("ExpectInsertOneTracksDocumentAndOptions", func(t *testing.T) { + mock := NewMockDatabase().ExpectInsertOne("inserted-id", nil) + document := bson.M{"name": "organisation"} + opts := options.InsertOne().SetBypassDocumentValidation(true) + + result, err := mock.InsertOne(context.Background(), "testdb", "organisation", document, opts) + if err != nil { + t.Fatalf("InsertOne error: %v", err) + } + if result != "inserted-id" { + t.Fatalf("InsertOne result = %#v", result) + } + if len(mock.InsertOneCalls) != 1 { + t.Fatalf("InsertOne calls = %d", len(mock.InsertOneCalls)) + } + call := mock.InsertOneCalls[0] + if call.Db != "testdb" || call.Collection != "organisation" || len(call.Opts) != 1 || call.Opts[0] != opts { + t.Fatalf("unexpected InsertOne call: %#v", call) + } + }) + + t.Run("QueueInsertManyPreservesResponsesAndReset", func(t *testing.T) { + expectedErr := errors.New("insert many failed") + mock := NewMockDatabase(). + QueueInsertMany([]any{"first", "second"}, nil). + QueueInsertMany(nil, expectedErr) + documents := []any{bson.M{"name": "first"}, bson.M{"name": "second"}} + opts := options.InsertMany().SetOrdered(false) + + result, err := mock.InsertMany(context.Background(), "testdb", "organisation", documents, opts) + if err != nil { + t.Fatalf("first InsertMany error: %v", err) + } + insertedIDs, ok := result.([]any) + if !ok || len(insertedIDs) != 2 { + t.Fatalf("first InsertMany result = %#v", result) + } + + result, err = mock.InsertMany(context.Background(), "testdb", "organisation", documents) + if !errors.Is(err, expectedErr) || result != nil { + t.Fatalf("second InsertMany result = %#v, error = %v", result, err) + } + if len(mock.InsertManyCalls) != 2 || len(mock.InsertManyCalls[0].Opts) != 1 || mock.InsertManyCalls[0].Opts[0] != opts { + t.Fatalf("InsertMany calls = %#v", mock.InsertManyCalls) + } + + mock.Reset() + if len(mock.InsertManyCalls) != 0 || len(mock.InsertManyQueue) != 0 { + t.Fatal("expected InsertMany state to be cleared") + } + }) + + t.Run("ExpectInsertManyReturnsConfiguredResult", func(t *testing.T) { + mock := NewMockDatabase().ExpectInsertMany([]any{"first", "second"}, nil) + + result, err := mock.InsertMany(context.Background(), "testdb", "organisation", []any{bson.M{}, bson.M{}}) + if err != nil { + t.Fatalf("InsertMany error: %v", err) + } + insertedIDs, ok := result.([]any) + if !ok || len(insertedIDs) != 2 || insertedIDs[0] != "first" || insertedIDs[1] != "second" { + t.Fatalf("InsertMany result = %#v", result) + } + }) + + t.Run("QueueInsertOnePreservesErrorAndReset", func(t *testing.T) { + expectedErr := errors.New("insert one failed") + mock := NewMockDatabase().QueueInsertOne(nil, expectedErr) + + result, err := mock.InsertOne(context.Background(), "testdb", "organisation", bson.M{}) + if !errors.Is(err, expectedErr) || result != nil { + t.Fatalf("InsertOne result = %#v, error = %v", result, err) + } + if len(mock.InsertOneCalls) != 1 { + t.Fatalf("InsertOne calls = %d", len(mock.InsertOneCalls)) + } + + mock.Reset() + if len(mock.InsertOneCalls) != 0 || len(mock.InsertOneQueue) != 0 { + t.Fatal("expected InsertOne state to be cleared") + } + }) +} diff --git a/pkg/database/mongodb.go b/pkg/database/mongodb.go index d2881bd..f8b7051 100644 --- a/pkg/database/mongodb.go +++ b/pkg/database/mongodb.go @@ -487,11 +487,12 @@ func (m *MongoClient) FindOneAndUpdate(ctx context.Context, db string, collectio return &SingleResult{result: coll.FindOneAndUpdate(ctx, filter, update, findOneAndUpdateOpts...)} } -// InsertOne inserts a single document into the specified database and collection +// InsertOne inserts a single document into the specified database and collection. +// Supports *moptions.InsertOneOptions in opts. func (m *MongoClient) InsertOne(ctx context.Context, db string, collection string, document any, opts ...any) (any, error) { coll := m.Client.Database(db).Collection(collection) - result, err := coll.InsertOne(ctx, document) + result, err := coll.InsertOne(ctx, document, insertOneOptions(opts...)...) if err != nil { return nil, err } @@ -499,11 +500,24 @@ func (m *MongoClient) InsertOne(ctx context.Context, db string, collection strin return result.InsertedID, nil } -// InsertMany inserts multiple documents into the specified database and collection +// insertOneOptions retains supported MongoDB options and ignores unrelated +// wrapper options, consistent with the other DatabaseInterface methods. +func insertOneOptions(opts ...any) []*moptions.InsertOneOptions { + var insertOpts []*moptions.InsertOneOptions + for _, opt := range opts { + if io, ok := opt.(*moptions.InsertOneOptions); ok { + insertOpts = append(insertOpts, io) + } + } + return insertOpts +} + +// InsertMany inserts multiple documents into the specified database and collection. +// Supports *moptions.InsertManyOptions in opts. func (m *MongoClient) InsertMany(ctx context.Context, db string, collection string, documents []any, opts ...any) (any, error) { coll := m.Client.Database(db).Collection(collection) - result, err := coll.InsertMany(ctx, documents) + result, err := coll.InsertMany(ctx, documents, insertManyOptions(opts...)...) if err != nil { return nil, err } @@ -511,6 +525,18 @@ func (m *MongoClient) InsertMany(ctx context.Context, db string, collection stri return result.InsertedIDs, nil } +// insertManyOptions retains supported MongoDB options and ignores unrelated +// wrapper options, consistent with the other DatabaseInterface methods. +func insertManyOptions(opts ...any) []*moptions.InsertManyOptions { + var insertOpts []*moptions.InsertManyOptions + for _, opt := range opts { + if io, ok := opt.(*moptions.InsertManyOptions); ok { + insertOpts = append(insertOpts, io) + } + } + return insertOpts +} + // SingleResult wraps a MongoDB single result for fluent API usage type SingleResult struct { result *mongo.SingleResult diff --git a/pkg/database/mongodb_test.go b/pkg/database/mongodb_test.go index 28aad1c..afe7533 100644 --- a/pkg/database/mongodb_test.go +++ b/pkg/database/mongodb_test.go @@ -13,8 +13,31 @@ import ( "github.com/uug-ai/models/pkg/models" "go.mongodb.org/mongo-driver/bson" + moptions "go.mongodb.org/mongo-driver/mongo/options" ) +func TestInsertOptions(t *testing.T) { + t.Run("InsertOneRetainsSupportedOptions", func(t *testing.T) { + first := moptions.InsertOne().SetBypassDocumentValidation(true) + second := moptions.InsertOne().SetComment("insert-one") + + got := insertOneOptions(first, moptions.Find(), second) + if len(got) != 2 || got[0] != first || got[1] != second { + t.Fatalf("insert one options = %#v", got) + } + }) + + t.Run("InsertManyRetainsSupportedOptions", func(t *testing.T) { + first := moptions.InsertMany().SetOrdered(false) + second := moptions.InsertMany().SetBypassDocumentValidation(true) + + got := insertManyOptions(first, moptions.FindOne(), second) + if len(got) != 2 || got[0] != first || got[1] != second { + t.Fatalf("insert many options = %#v", got) + } + }) +} + func envVarsSet(keys ...string) bool { for _, key := range keys { if os.Getenv(key) == "" {