Skip to content
2 changes: 1 addition & 1 deletion exporter/exporterhelper/internal/queuebatch/batcher.go
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ func NewBatcher(cfg configoptional.Optional[BatchConfig], set batcherSettings[re
}

if set.partitioner == nil {
return newPartitionBatcher(*cfg.Get(), sizer, set.mergeCtx, newWorkerPool(set.maxWorkers), set.next, set.logger, nil), nil
return newShardedBatcher(*cfg.Get(), sizer, set.mergeCtx, newWorkerPool(set.maxWorkers), set.next, set.logger, 1)
}

mb, err := newMultiBatcher(*cfg.Get(), sizer, newWorkerPool(set.maxWorkers), set.partitioner, set.mergeCtx, set.next, set.logger)
Expand Down
70 changes: 70 additions & 0 deletions exporter/exporterhelper/internal/queuebatch/sharded_batcher.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
// Copyright The OpenTelemetry Authors
// SPDX-License-Identifier: Apache-2.0

package queuebatch // import "go.opentelemetry.io/collector/exporter/exporterhelper/internal/queuebatch"

import (
"context"
"fmt"
"sync/atomic"

"go.uber.org/multierr"
"go.uber.org/zap"

"go.opentelemetry.io/collector/component"
"go.opentelemetry.io/collector/exporter/exporterhelper/internal/queue"
"go.opentelemetry.io/collector/exporter/exporterhelper/internal/request"
"go.opentelemetry.io/collector/exporter/exporterhelper/internal/sender"
)

var _ Batcher[request.Request] = (*shardedBatcher)(nil)

// shardedBatcher spreads unpartitioned requests across independent partitionBatcher instances.
// All shards share one worker pool, so worker concurrency remains bounded by the queue setting.
type shardedBatcher struct {
shards []*partitionBatcher
next atomic.Uint64
}

func newShardedBatcher(
cfg BatchConfig,
sizer request.Sizer,
mergeCtx func(context.Context, context.Context) context.Context,
wp *workerPool,
next sender.SendFunc[request.Request],
logger *zap.Logger,
shardCount int,
) (*shardedBatcher, error) {
if shardCount <= 0 {
return nil, fmt.Errorf("queue_batch: shard count must be positive, found %d", shardCount)
}

sb := &shardedBatcher{
shards: make([]*partitionBatcher, 0, shardCount),
}
for range shardCount {
sb.shards = append(sb.shards, newPartitionBatcher(cfg, sizer, mergeCtx, wp, next, logger, nil))
}
return sb, nil
}

func (sb *shardedBatcher) Start(ctx context.Context, host component.Host) error {
var err error
for _, shard := range sb.shards {
err = multierr.Append(err, shard.Start(ctx, host))
}
return err
}

func (sb *shardedBatcher) Consume(ctx context.Context, req request.Request, done queue.Done) {
idx := sb.next.Add(1) - 1
sb.shards[idx%uint64(len(sb.shards))].Consume(ctx, req, done)
}

func (sb *shardedBatcher) Shutdown(ctx context.Context) error {
var err error
for _, shard := range sb.shards {
err = multierr.Append(err, shard.Shutdown(ctx))
}
return err
}
162 changes: 162 additions & 0 deletions exporter/exporterhelper/internal/queuebatch/sharded_batcher_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
// Copyright The OpenTelemetry Authors
// SPDX-License-Identifier: Apache-2.0

package queuebatch

import (
"context"
"runtime"
"testing"
"time"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.uber.org/zap"

"go.opentelemetry.io/collector/component/componenttest"
"go.opentelemetry.io/collector/config/configoptional"
"go.opentelemetry.io/collector/exporter/exporterhelper/internal/request"
"go.opentelemetry.io/collector/exporter/exporterhelper/internal/requesttest"
)

func TestShardedBatcherRequiresPositiveShardCount(t *testing.T) {
sink := requesttest.NewSink()
sb, err := newShardedBatcher(BatchConfig{}, request.NewItemsSizer(), nil, newWorkerPool(1), sink.Export, zap.NewNop(), 0)

require.Error(t, err)
assert.Nil(t, sb)
}

func TestNewBatcherUsesOneShardForUnpartitionedBatching(t *testing.T) {
sink := requesttest.NewSink()
batcher, err := NewBatcher(configoptional.Some(BatchConfig{
FlushTimeout: 50 * time.Millisecond,
Sizer: request.SizerTypeItems,
MinSize: 10,
}), batcherSettings[request.Request]{
next: sink.Export,
maxWorkers: 1,
logger: zap.NewNop(),
})

require.NoError(t, err)
sb, ok := batcher.(*shardedBatcher)
require.True(t, ok)
assert.Len(t, sb.shards, 1)
}

func TestShardedBatcherOneShardFlushesOnMinSize(t *testing.T) {
sink := requesttest.NewSink()
sb, err := newShardedBatcher(BatchConfig{
FlushTimeout: 0,
Sizer: request.SizerTypeItems,
MinSize: 10,
}, request.NewItemsSizer(), nil, newWorkerPool(1), sink.Export, zap.NewNop(), 1)
require.NoError(t, err)
require.NoError(t, sb.Start(context.Background(), componenttest.NewNopHost()))
t.Cleanup(func() {
require.NoError(t, sb.Shutdown(context.Background()))
})

done := newFakeDone()
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 4}, done)
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 6}, done)

assert.Eventually(t, func() bool {
return sink.RequestsCount() == 1 && sink.ItemsCount() == 10 &&
done.success.Load() == 2 && done.errors.Load() == 0
}, time.Second, 10*time.Millisecond)
}

func TestShardedBatcherMultipleShardsFlushOnMinSize(t *testing.T) {
sink := requesttest.NewSink()
sb, err := newShardedBatcher(BatchConfig{
FlushTimeout: 0,
Sizer: request.SizerTypeItems,
MinSize: 10,
}, request.NewItemsSizer(), nil, newWorkerPool(2), sink.Export, zap.NewNop(), 2)
require.NoError(t, err)
require.NoError(t, sb.Start(context.Background(), componenttest.NewNopHost()))
t.Cleanup(func() {
require.NoError(t, sb.Shutdown(context.Background()))
})

done := newFakeDone()
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 4}, done)
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 4}, done)
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 6}, done)
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 6}, done)

assert.Eventually(t, func() bool {
return sink.RequestsCount() == 2 && sink.ItemsCount() == 20 &&
done.success.Load() == 4 && done.errors.Load() == 0
}, time.Second, 10*time.Millisecond)
}

func TestShardedBatcherOneShardFlushesOnTimeout(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("Skipping test on Windows, see https://github.com/open-telemetry/opentelemetry-collector/issues/11869")
}

sink := requesttest.NewSink()
sb, err := newShardedBatcher(BatchConfig{
FlushTimeout: 50 * time.Millisecond,
Sizer: request.SizerTypeItems,
MinSize: 10,
}, request.NewItemsSizer(), nil, newWorkerPool(1), sink.Export, zap.NewNop(), 1)
require.NoError(t, err)
require.NoError(t, sb.Start(context.Background(), componenttest.NewNopHost()))
t.Cleanup(func() {
require.NoError(t, sb.Shutdown(context.Background()))
})

done := newFakeDone()
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 4}, done)

assert.Eventually(t, func() bool {
return sink.RequestsCount() == 1 && sink.ItemsCount() == 4 &&
done.success.Load() == 1 && done.errors.Load() == 0
}, time.Second, 10*time.Millisecond)
}

func TestShardedBatcherOneShardDrainsOnShutdown(t *testing.T) {
sink := requesttest.NewSink()
sb, err := newShardedBatcher(BatchConfig{
FlushTimeout: 0,
Sizer: request.SizerTypeItems,
MinSize: 10,
}, request.NewItemsSizer(), nil, newWorkerPool(1), sink.Export, zap.NewNop(), 1)
require.NoError(t, err)
require.NoError(t, sb.Start(context.Background(), componenttest.NewNopHost()))

done := newFakeDone()
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 4}, done)
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 3}, done)

require.NoError(t, sb.Shutdown(context.Background()))
assert.Equal(t, 1, sink.RequestsCount())
assert.Equal(t, 7, sink.ItemsCount())
assert.EqualValues(t, 2, done.success.Load())
assert.EqualValues(t, 0, done.errors.Load())
}

func TestShardedBatcherMultipleShardsDrainOnShutdown(t *testing.T) {
sink := requesttest.NewSink()
sb, err := newShardedBatcher(BatchConfig{
FlushTimeout: 0,
Sizer: request.SizerTypeItems,
MinSize: 10,
}, request.NewItemsSizer(), nil, newWorkerPool(2), sink.Export, zap.NewNop(), 2)
require.NoError(t, err)
require.NoError(t, sb.Start(context.Background(), componenttest.NewNopHost()))

done := newFakeDone()
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 4}, done)
sb.Consume(context.Background(), &requesttest.FakeRequest{Items: 3}, done)

require.NoError(t, sb.Shutdown(context.Background()))
assert.Equal(t, 2, sink.RequestsCount())
assert.Equal(t, 7, sink.ItemsCount())
assert.EqualValues(t, 2, done.success.Load())
assert.EqualValues(t, 0, done.errors.Load())
}
Loading