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/aggregate.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
package database

import "context"

// Aggregator is the optional database capability for executing MongoDB
// aggregation pipelines. Pipeline construction and result decoding remain
// consumer concerns.
type Aggregator interface {
Aggregate(ctx context.Context, db string, collection string, pipeline any, opts ...any) FindResultInterface
}
57 changes: 57 additions & 0 deletions pkg/database/mock.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,9 @@ type MockDatabase struct {
// CreateIndexesFunc allows customizing index creation behavior.
CreateIndexesFunc func(ctx context.Context, db string, collection string, indexes []mongo.IndexModel, opts ...any) ([]string, error)

// AggregateFunc allows customizing Aggregate behavior.
AggregateFunc func(ctx context.Context, db string, collection string, pipeline any, opts ...any) FindResultInterface

// UpdateOneFunc allows customizing UpdateOne behavior
UpdateOneFunc func(ctx context.Context, db string, collection string, filter any, update any, opts ...any) (UpdateResultInterface, error)

Expand Down Expand Up @@ -58,6 +61,7 @@ type MockDatabase struct {
CountQueue []CountResponse
InsertOneQueue []InsertResponse
InsertManyQueue []InsertResponse
AggregateQueue []FindResponse

// Call tracking
PingCalls []PingCall
Expand All @@ -71,10 +75,12 @@ type MockDatabase struct {
CountCalls []CountCall
InsertOneCalls []InsertOneCall
InsertManyCalls []InsertManyCall
AggregateCalls []AggregateCall
}

var _ DatabaseInterface = (*MockDatabase)(nil)
var _ IndexManager = (*MockDatabase)(nil)
var _ Aggregator = (*MockDatabase)(nil)

// MockSingleResult implements SingleResultInterface for testing
type MockSingleResult struct {
Expand Down Expand Up @@ -286,6 +292,15 @@ type CreateIndexesCall struct {
Opts []any
}

// AggregateCall records an aggregation call.
type AggregateCall struct {
Ctx context.Context
Db string
Collection string
Pipeline any
Opts []any
}

// UpdateOneCall records a call to UpdateOne
type UpdateOneCall struct {
Ctx context.Context
Expand Down Expand Up @@ -362,6 +377,9 @@ func NewMockDatabase() *MockDatabase {
CreateIndexesFunc: func(ctx context.Context, db string, collection string, indexes []mongo.IndexModel, opts ...any) ([]string, error) {
return nil, nil
},
AggregateFunc: func(ctx context.Context, db string, collection string, pipeline any, opts ...any) FindResultInterface {
return &MockFindResult{results: []any{}, err: nil}
},
UpdateOneFunc: func(ctx context.Context, db string, collection string, filter any, update any, opts ...any) (UpdateResultInterface, error) {
return &MockUpdateResult{matchedCount: 1, modifiedCount: 1}, nil
},
Expand All @@ -385,6 +403,7 @@ func NewMockDatabase() *MockDatabase {
CountCalls: []CountCall{},
InsertOneCalls: []InsertOneCall{},
InsertManyCalls: []InsertManyCall{},
AggregateCalls: []AggregateCall{},
PingQueue: []PingResponse{},
FindQueue: []FindResponse{},
FindOneQueue: []FindOneResponse{},
Expand All @@ -396,6 +415,7 @@ func NewMockDatabase() *MockDatabase {
CountQueue: []CountResponse{},
InsertOneQueue: []InsertResponse{},
InsertManyQueue: []InsertResponse{},
AggregateQueue: []FindResponse{},
}
}

Expand Down Expand Up @@ -467,6 +487,27 @@ func (m *MockDatabase) Find(ctx context.Context, db string, collection string, f
return &MockFindResult{results: result, err: err}
}

// Aggregate implements Aggregator.
func (m *MockDatabase) Aggregate(ctx context.Context, db string, collection string, pipeline any, opts ...any) FindResultInterface {
m.AggregateCalls = append(m.AggregateCalls, AggregateCall{
Ctx: ctx,
Db: db,
Collection: collection,
Pipeline: pipeline,
Opts: opts,
})

if len(m.AggregateQueue) > 0 {
response := m.AggregateQueue[0]
m.AggregateQueue = m.AggregateQueue[1:]
return &MockFindResult{results: response.Result, err: response.Err}
}
if m.AggregateFunc != nil {
return m.AggregateFunc(ctx, db, collection, pipeline, opts...)
}
return &MockFindResult{results: []any{}, err: nil}
}

// FindOne implements DatabaseInterface
func (m *MockDatabase) FindOne(ctx context.Context, db string, collection string, filter any, opts ...any) SingleResultInterface {
m.FindOneCalls = append(m.FindOneCalls, FindOneCall{
Expand Down Expand Up @@ -804,6 +845,7 @@ func (m *MockDatabase) Reset() {
m.CountCalls = []CountCall{}
m.InsertOneCalls = []InsertOneCall{}
m.InsertManyCalls = []InsertManyCall{}
m.AggregateCalls = []AggregateCall{}
m.PingQueue = []PingResponse{}
m.FindQueue = []FindResponse{}
m.FindOneQueue = []FindOneResponse{}
Expand All @@ -815,6 +857,7 @@ func (m *MockDatabase) Reset() {
m.CountQueue = []CountResponse{}
m.InsertOneQueue = []InsertResponse{}
m.InsertManyQueue = []InsertResponse{}
m.AggregateQueue = []FindResponse{}
}

// ExpectPing sets up an expectation for Ping
Expand Down Expand Up @@ -857,6 +900,14 @@ func (m *MockDatabase) ExpectCreateIndexes(names []string, err error) *MockDatab
return m
}

// ExpectAggregate sets up an aggregation response.
func (m *MockDatabase) ExpectAggregate(result any, err error) *MockDatabase {
m.AggregateFunc = func(ctx context.Context, db string, collection string, pipeline any, opts ...any) FindResultInterface {
return &MockFindResult{results: result, err: err}
}
return m
}

// ExpectUpdateOne sets up an expectation for UpdateOne
func (m *MockDatabase) ExpectUpdateOne(matchedCount, modifiedCount, upsertedCount int64, upsertedID any, err error) *MockDatabase {
m.UpdateOneFunc = func(ctx context.Context, db string, collection string, filter any, update any, opts ...any) (UpdateResultInterface, error) {
Expand Down Expand Up @@ -925,6 +976,12 @@ func (m *MockDatabase) QueueCreateIndexes(names []string, err error) *MockDataba
return m
}

// QueueAggregate adds an aggregation response to the queue.
func (m *MockDatabase) QueueAggregate(result any, err error) *MockDatabase {
m.AggregateQueue = append(m.AggregateQueue, FindResponse{Result: result, Err: err})
return m
}

// QueueUpdateOne adds an UpdateOne response to the queue for sequential calls
func (m *MockDatabase) QueueUpdateOne(matchedCount, modifiedCount, upsertedCount int64, upsertedID any, err error) *MockDatabase {
m.UpdateOneQueue = append(m.UpdateOneQueue, UpdateOneResponse{
Expand Down
42 changes: 42 additions & 0 deletions pkg/database/mock_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -911,3 +911,45 @@ func TestMockDatabaseInsert(t *testing.T) {
}
})
}

func TestMockDatabaseAggregate(t *testing.T) {
t.Run("ExpectAggregateTracksPipelineAndOptions", func(t *testing.T) {
mock := NewMockDatabase().ExpectAggregate([]bson.M{{"count": 2}}, nil)
pipeline := mongo.Pipeline{{{Key: "$group", Value: bson.D{{Key: "_id", Value: "$status"}}}}}
opts := options.Aggregate().SetAllowDiskUse(true)

var result []bson.M
err := mock.Aggregate(context.Background(), "testdb", "events", pipeline, opts).All(&result)
if err != nil {
t.Fatalf("Aggregate error: %v", err)
}
if len(result) != 1 || result[0]["count"] != float64(2) {
t.Fatalf("Aggregate result = %#v", result)
}
if len(mock.AggregateCalls) != 1 {
t.Fatalf("Aggregate calls = %d", len(mock.AggregateCalls))
}
call := mock.AggregateCalls[0]
if call.Db != "testdb" || call.Collection != "events" || len(call.Opts) != 1 || call.Opts[0] != opts {
t.Fatalf("unexpected Aggregate call: %#v", call)
}
})

t.Run("QueueAggregatePreservesErrorAndReset", func(t *testing.T) {
expectedErr := errors.New("aggregate failed")
mock := NewMockDatabase().QueueAggregate(nil, expectedErr)

result := mock.Aggregate(context.Background(), "testdb", "events", mongo.Pipeline{})
if !errors.Is(result.Err(), expectedErr) {
t.Fatalf("Aggregate error = %v, want %v", result.Err(), expectedErr)
}
if len(mock.AggregateCalls) != 1 {
t.Fatalf("Aggregate calls = %d", len(mock.AggregateCalls))
}

mock.Reset()
if len(mock.AggregateCalls) != 0 || len(mock.AggregateQueue) != 0 {
t.Fatal("expected Aggregate state to be cleared")
}
})
}
23 changes: 23 additions & 0 deletions pkg/database/mongodb.go
Original file line number Diff line number Diff line change
Expand Up @@ -231,6 +231,7 @@ type MongoClient struct {

var _ DatabaseInterface = (*MongoClient)(nil)
var _ IndexManager = (*MongoClient)(nil)
var _ Aggregator = (*MongoClient)(nil)

// NewMongoClient creates a new MongoClient with the provided MongoDB settings
func NewMongoClient(options *MongoOptions) (DatabaseInterface, error) {
Expand Down Expand Up @@ -450,6 +451,28 @@ func (m *MongoClient) Find(ctx context.Context, db string, collection string, fi
return &FindResult{cursor: cursor, ctx: ctx, err: err}
}

// Aggregate executes an aggregation pipeline on the specified collection.
// Results support both slice-wide decoding and optional cursor iteration.
// Supports *moptions.AggregateOptions in opts.
func (m *MongoClient) Aggregate(ctx context.Context, db string, collection string, pipeline any, opts ...any) FindResultInterface {
coll := m.Client.Database(db).Collection(collection)

cursor, err := coll.Aggregate(ctx, pipeline, aggregateOptions(opts...)...)
return &FindResult{cursor: cursor, ctx: ctx, err: err}
}

// aggregateOptions retains supported MongoDB options and ignores unrelated
// wrapper options, consistent with the other optional capabilities.
func aggregateOptions(opts ...any) []*moptions.AggregateOptions {
var aggregateOpts []*moptions.AggregateOptions
for _, opt := range opts {
if ao, ok := opt.(*moptions.AggregateOptions); ok {
aggregateOpts = append(aggregateOpts, ao)
}
}
return aggregateOpts
}

// FindOne executes a findOne query on the specified database and collection.
// Returns a SingleResult that can be used with .Into() or .Raw() for fluent decoding.
// Supports *moptions.FindOneOptions and *Projection in opts.
Expand Down
10 changes: 10 additions & 0 deletions pkg/database/mongodb_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,16 @@ func TestInsertOptions(t *testing.T) {
})
}

func TestAggregateOptions(t *testing.T) {
first := moptions.Aggregate().SetAllowDiskUse(true)
second := moptions.Aggregate().SetBatchSize(100)

got := aggregateOptions(first, moptions.Find(), second)
if len(got) != 2 || got[0] != first || got[1] != second {
t.Fatalf("aggregate options = %#v", got)
}
}

func envVarsSet(keys ...string) bool {
for _, key := range keys {
if os.Getenv(key) == "" {
Expand Down
Loading