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
102 changes: 102 additions & 0 deletions pkg/database/mock.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -50,6 +56,8 @@ type MockDatabase struct {
DeleteOneQueue []DeleteResponse
DeleteManyQueue []DeleteResponse
CountQueue []CountResponse
InsertOneQueue []InsertResponse
InsertManyQueue []InsertResponse

// Call tracking
PingCalls []PingCall
Expand All @@ -61,6 +69,8 @@ type MockDatabase struct {
DeleteOneCalls []DeleteCall
DeleteManyCalls []DeleteCall
CountCalls []CountCall
InsertOneCalls []InsertOneCall
InsertManyCalls []InsertManyCall
}

var _ DatabaseInterface = (*MockDatabase)(nil)
Expand Down Expand Up @@ -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
Expand All @@ -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{
Expand Down Expand Up @@ -349,6 +383,8 @@ func NewMockDatabase() *MockDatabase {
DeleteOneCalls: []DeleteCall{},
DeleteManyCalls: []DeleteCall{},
CountCalls: []CountCall{},
InsertOneCalls: []InsertOneCall{},
InsertManyCalls: []InsertManyCall{},
PingQueue: []PingResponse{},
FindQueue: []FindResponse{},
FindOneQueue: []FindOneResponse{},
Expand All @@ -358,6 +394,8 @@ func NewMockDatabase() *MockDatabase {
DeleteOneQueue: []DeleteResponse{},
DeleteManyQueue: []DeleteResponse{},
CountQueue: []CountResponse{},
InsertOneQueue: []InsertResponse{},
InsertManyQueue: []InsertResponse{},
}
}

Expand Down Expand Up @@ -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")
}

Expand All @@ -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{}
Expand All @@ -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
Expand Down Expand Up @@ -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
}
85 changes: 85 additions & 0 deletions pkg/database/mock_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
})
}
34 changes: 30 additions & 4 deletions pkg/database/mongodb.go
Original file line number Diff line number Diff line change
Expand Up @@ -487,30 +487,56 @@ 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
}

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
}

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
Expand Down
23 changes: 23 additions & 0 deletions pkg/database/mongodb_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) == "" {
Expand Down
Loading