diff --git a/agent/state.go b/agent/state.go index 657be4f8..b1da510c 100644 --- a/agent/state.go +++ b/agent/state.go @@ -41,6 +41,7 @@ type StateMachine struct { Transitions map[state]map[event]state StateFunctions map[state]func() logger *slog.Logger + wg *sync.WaitGroup } // NewStateMachine creates a new StateMachine. @@ -51,6 +52,7 @@ func NewStateMachine(logger *slog.Logger) *StateMachine { Transitions: make(map[state]map[event]state), StateFunctions: make(map[state]func()), logger: logger, + wg: &sync.WaitGroup{}, } sm.Transitions[idle] = make(map[event]state) @@ -76,6 +78,8 @@ func NewStateMachine(logger *slog.Logger) *StateMachine { // Start the state machine. func (sm *StateMachine) Start(ctx context.Context) { + sm.wg.Add(1) + defer sm.wg.Done() for { select { case event := <-sm.EventChan: diff --git a/agent/state_test.go b/agent/state_test.go index d20e44c8..a1e38bc4 100644 --- a/agent/state_test.go +++ b/agent/state_test.go @@ -27,12 +27,11 @@ func TestStateMachineTransitions(t *testing.T) { for _, testCase := range testCases { t.Run(fmt.Sprintf("Transition from %v to %v", testCase.fromState, testCase.expected), func(t *testing.T) { sm := NewStateMachine(mglog.NewMock()) - done := make(chan struct{}) ctx, cancel := context.WithCancel(context.Background()) go func() { sm.Start(ctx) - close(done) }() + sm.wg.Wait() sm.SetState(testCase.fromState) sm.SendEvent(testCase.event) @@ -42,7 +41,6 @@ func TestStateMachineTransitions(t *testing.T) { } close(sm.EventChan) cancel() - <-done }) } }