diff --git a/lib/utils/batcher/batcher.go b/lib/utils/batcher/batcher.go index f0be686d1ee..c5e87b6490f 100644 --- a/lib/utils/batcher/batcher.go +++ b/lib/utils/batcher/batcher.go @@ -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 { diff --git a/lib/utils/batcher/batcher_test.go b/lib/utils/batcher/batcher_test.go index 05720188708..323a170b144 100644 --- a/lib/utils/batcher/batcher_test.go +++ b/lib/utils/batcher/batcher_test.go @@ -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) diff --git a/lib/utils/batcher/state.go b/lib/utils/batcher/state.go new file mode 100644 index 00000000000..afaaa64f337 --- /dev/null +++ b/lib/utils/batcher/state.go @@ -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 . + +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 +} diff --git a/lib/utils/batcher/state_test.go b/lib/utils/batcher/state_test.go new file mode 100644 index 00000000000..f7303a184ae --- /dev/null +++ b/lib/utils/batcher/state_test.go @@ -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 . +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) +}