Add utils/batcher state (#59457)

This commit is contained in:
Marek Smoliński
2025-09-23 15:52:01 +00:00
committed by GitHub
parent 6d13965820
commit 667ca241a4
4 changed files with 214 additions and 12 deletions
+15 -1
View File
@@ -56,6 +56,21 @@ func Run[T any](ctx context.Context, events <-chan T, fn func(batch []T) error,
return trace.Wrap(collector.Run(ctx, events, fn))
}
// RunWithState is like Run but also provides the current State to the processing function.
// The state is managed by the provided StateMonitor and updated based on the batch size.
func RunWithState[T any](
ctx context.Context,
events <-chan T,
fn func(batch []T, state State) error,
stateManager *StateMonitor,
opts ...Option,
) error {
return trace.Wrap(Run(ctx, events, func(batch []T) error {
currentState := stateManager.UpdateState(len(batch))
return trace.Wrap(fn(batch, currentState))
}, opts...))
}
// Iter allows to iterate over batches of events.
func Iter[T any](ctx context.Context, events <-chan T, opts ...Option) iter.Seq2[[]T, error] {
collector := New[T](opts...)
@@ -82,7 +97,6 @@ func WithWindow(d time.Duration) Option {
}
// WithThreshold sets the maximum number of events per batch.
// Default is unlimited (bounded only by window duration).
func WithThreshold(n int) Option {
return func(c *config) {
if n > 0 {
+11 -11
View File
@@ -26,15 +26,15 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
. "github.com/gravitational/teleport/lib/utils/batcher"
"github.com/gravitational/teleport/lib/utils/batcher"
)
func TestCollectBatch_TimeWindow(t *testing.T) {
fakeClock := clockwork.NewFakeClock()
collector := New[string](
WithWindow(100*time.Millisecond),
WithClock(fakeClock),
WithThreshold(10),
collector := batcher.New[string](
batcher.WithWindow(100*time.Millisecond),
batcher.WithClock(fakeClock),
batcher.WithThreshold(10),
)
events := make(chan string)
@@ -63,7 +63,7 @@ func TestCollectBatch_TimeWindow(t *testing.T) {
}
func TestCollectBatch_Threshold(t *testing.T) {
collector := New[int](WithThreshold(3))
collector := batcher.New[int](batcher.WithThreshold(3))
events := make(chan int, 10)
for i := 1; i <= 7; i++ {
@@ -83,7 +83,7 @@ func TestCollectBatch_Threshold(t *testing.T) {
}
func TestCollectBatch_ChannelClosed(t *testing.T) {
collector := New[string]()
collector := batcher.New[string]()
events := make(chan string, 2)
events <- "event1"
@@ -99,7 +99,7 @@ func TestCollectBatch_ChannelClosed(t *testing.T) {
}
func TestCollectBatch_ContextCanceled(t *testing.T) {
collector := New[string](WithWindow(1 * time.Hour))
collector := batcher.New[string](batcher.WithWindow(1 * time.Hour))
events := make(chan string, 2)
events <- "event1"
@@ -113,9 +113,9 @@ func TestCollectBatch_ContextCanceled(t *testing.T) {
func TestCollectBatch_EmptyChannel(t *testing.T) {
fakeClock := clockwork.NewFakeClock()
collector := New[string](
WithWindow(100*time.Millisecond),
WithClock(fakeClock),
collector := batcher.New[string](
batcher.WithWindow(100*time.Millisecond),
batcher.WithClock(fakeClock),
)
events := make(chan string)
+123
View File
@@ -0,0 +1,123 @@
// Teleport
// Copyright (C) 2025 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package batcher
import "sync"
// State represents the current processing mode based on batch volume.
type State int
const (
// StateNormal indicates regular processing mode with normal event volume.
StateNormal State = iota
// StateOverloaded indicates high-volume processing mode when batch sizes
// when threshold is reached.
StateOverloaded
)
// String returns a human-readable representation of the State.
func (s State) String() string {
switch s {
case StateNormal:
return "Normal"
case StateOverloaded:
return "Overloaded"
default:
return "Unknown"
}
}
// StateConfig configures state transitions and callbacks for a StateMonitor.
type StateConfig struct {
// OnEnterOverloaded is called when transitioning from Normal to Overloaded state.
// This callback is invoked when batch size reaches or exceeds the threshold.
// Is not provided, no action is taken on state transition.
OnEnterOverloaded func()
// OnExitOverloaded is called when transitioning from Overloaded to Normal state.
// This callback is invoked when batch size drops below threshold/2.
//
// WARN: For simplicity there is not timer-based transition back to Normal state,
// it only happens when batch size drops below threshold/2.
// So if the last batch is larger and there are no new events, the state will remain Overloaded.
// Is not provided, no action is taken on state transition.
OnExitOverloaded func()
// Threshold is the batch size threshold for state transitions.
// Normal -> Overloaded: when batchSize >= Threshold
// Overloaded -> Normal: when batchSize < Threshold/2
Threshold int
}
// StateMonitor monitors batch sizes and manages state transitions based on configured thresholds.
type StateMonitor struct {
StateConfig
mu sync.RWMutex
state State
}
// NewStateMonitor creates a new StateMonitor with the given configuration.
// The monitor starts in Normal state by default.
func NewStateMonitor(config StateConfig) *StateMonitor {
return &StateMonitor{
state: StateNormal,
StateConfig: config,
}
}
// GetState returns the current state of the monitor.
func (sm *StateMonitor) GetState() State {
sm.mu.RLock()
defer sm.mu.RUnlock()
return sm.state
}
// UpdateState evaluates the current batch size and updates the state if necessary.
func (sm *StateMonitor) UpdateState(batchSize int) State {
sm.mu.Lock()
oldState := sm.state
newState := sm.evaluateState(batchSize)
sm.state = newState
sm.mu.Unlock()
if oldState != newState {
switch {
case oldState != StateOverloaded && newState == StateOverloaded:
if sm.StateConfig.OnEnterOverloaded != nil {
sm.StateConfig.OnEnterOverloaded()
}
case oldState == StateOverloaded && newState == StateNormal:
if sm.StateConfig.OnExitOverloaded != nil {
sm.StateConfig.OnExitOverloaded()
}
}
}
return newState
}
func (sm *StateMonitor) evaluateState(batchSize int) State {
if batchSize >= sm.StateConfig.Threshold {
// Threshold is reached, enter Overloaded state.
return StateOverloaded
}
if batchSize < sm.StateConfig.Threshold/2 {
// The batch event size dropped below half the threshold, return to Normal state.
return StateNormal
}
return sm.state
}
+65
View File
@@ -0,0 +1,65 @@
// Teleport
// Copyright (C) 2025 Gravitational, Inc.
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU Affero General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
// GNU Affero General Public License for more details.
//
// You should have received a copy of the GNU Affero General Public License
// along with this program. If not, see <http://www.gnu.org/licenses/>.
package batcher_test
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/lib/utils/batcher"
)
func TestStateManage(t *testing.T) {
enterStormCount := 0
exitStormCount := 0
config := batcher.StateConfig{
Threshold: 100,
OnEnterOverloaded: func() {
enterStormCount++
},
OnExitOverloaded: func() {
exitStormCount++
},
}
sm := batcher.NewStateMonitor(config)
sm.UpdateState(50)
require.Equal(t, 0, enterStormCount)
require.Equal(t, 0, exitStormCount)
sm.UpdateState(100)
require.Equal(t, 1, enterStormCount)
require.Equal(t, 0, exitStormCount)
sm.UpdateState(150)
require.Equal(t, 1, enterStormCount)
require.Equal(t, 0, exitStormCount)
sm.UpdateState(40)
require.Equal(t, 1, enterStormCount)
require.Equal(t, 1, exitStormCount)
sm.UpdateState(20)
require.Equal(t, 1, enterStormCount)
require.Equal(t, 1, exitStormCount)
sm.UpdateState(100)
require.Equal(t, 2, enterStormCount)
require.Equal(t, 1, exitStormCount)
}