mirror of
https://github.com/lightninglabs/loop.git
synced 2026-08-13 12:33:03 +02:00
fsm: add ctx to SendEvent and Actions
This commit is contained in:
parent
b1719de5b9
commit
7b00baeac6
4 changed files with 39 additions and 25 deletions
|
|
@ -1,6 +1,7 @@
|
|||
package fsm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
|
|
@ -90,7 +91,9 @@ type InitStuffRequest struct {
|
|||
}
|
||||
|
||||
// initFSM is the action for the InitFSM state.
|
||||
func (e *ExampleFSM) initFSM(eventCtx EventContext) EventType {
|
||||
func (e *ExampleFSM) initFSM(_ context.Context, eventCtx EventContext,
|
||||
) EventType {
|
||||
|
||||
req, ok := eventCtx.(*InitStuffRequest)
|
||||
if !ok {
|
||||
return e.HandleError(
|
||||
|
|
@ -109,7 +112,9 @@ func (e *ExampleFSM) initFSM(eventCtx EventContext) EventType {
|
|||
}
|
||||
|
||||
// waitForStuff is an action that waits for stuff to happen.
|
||||
func (e *ExampleFSM) waitForStuff(eventCtx EventContext) EventType {
|
||||
func (e *ExampleFSM) waitForStuff(ctx context.Context, eventCtx EventContext,
|
||||
) EventType {
|
||||
|
||||
waitChan, err := e.service.WaitForStuffHappening()
|
||||
if err != nil {
|
||||
return e.HandleError(err)
|
||||
|
|
@ -117,7 +122,7 @@ func (e *ExampleFSM) waitForStuff(eventCtx EventContext) EventType {
|
|||
|
||||
go func() {
|
||||
<-waitChan
|
||||
err := e.SendEvent(OnStuffSuccess, nil)
|
||||
err := e.SendEvent(ctx, OnStuffSuccess, nil)
|
||||
if err != nil {
|
||||
log.Errorf("unable to send event: %v", err)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ func TestExampleFSM(t *testing.T) {
|
|||
tc := tc
|
||||
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctxb := context.Background()
|
||||
respondChan := make(chan string, 1)
|
||||
if req, ok := tc.eventCtx.(*InitStuffRequest); ok {
|
||||
req.respondChan = respondChan
|
||||
|
|
@ -102,7 +103,7 @@ func TestExampleFSM(t *testing.T) {
|
|||
exampleContext.RegisterObserver(cachedObserver)
|
||||
|
||||
err := exampleContext.SendEvent(
|
||||
tc.sendEvent, tc.eventCtx,
|
||||
ctxb, tc.sendEvent, tc.eventCtx,
|
||||
)
|
||||
require.Equal(t, tc.sendEventErr, err)
|
||||
|
||||
|
|
@ -195,6 +196,7 @@ func TestExampleFSMFlow(t *testing.T) {
|
|||
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
exampleContext, cachedObserver := getTestContext()
|
||||
ctxb := context.Background()
|
||||
|
||||
if tc.storeError != nil {
|
||||
exampleContext.store.(*mockStore).
|
||||
|
|
@ -208,8 +210,7 @@ func TestExampleFSMFlow(t *testing.T) {
|
|||
|
||||
go func() {
|
||||
err := exampleContext.SendEvent(
|
||||
OnRequestStuff,
|
||||
newInitStuffRequest(),
|
||||
ctxb, OnRequestStuff, newInitStuffRequest(),
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
|
|
@ -273,6 +274,7 @@ func TestObserverAsyncWait(t *testing.T) {
|
|||
service := &mockService{
|
||||
respondChan: make(chan bool),
|
||||
}
|
||||
ctxb := context.Background()
|
||||
|
||||
store := &mockStore{}
|
||||
|
||||
|
|
@ -282,7 +284,7 @@ func TestObserverAsyncWait(t *testing.T) {
|
|||
|
||||
t0 := time.Now()
|
||||
timeoutCtx, cancel := context.WithTimeout(
|
||||
context.Background(), tc.waitTime,
|
||||
ctxb, tc.waitTime,
|
||||
)
|
||||
defer cancel()
|
||||
|
||||
|
|
@ -293,8 +295,7 @@ func TestObserverAsyncWait(t *testing.T) {
|
|||
|
||||
go func() {
|
||||
err := exampleContext.SendEvent(
|
||||
OnRequestStuff,
|
||||
newInitStuffRequest(),
|
||||
ctxb, OnRequestStuff, newInitStuffRequest(),
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
|
|
|
|||
19
fsm/fsm.go
19
fsm/fsm.go
|
|
@ -1,6 +1,7 @@
|
|||
package fsm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
|
|
@ -45,7 +46,7 @@ type EventType string
|
|||
type EventContext interface{}
|
||||
|
||||
// Action represents the action to be executed in a given state.
|
||||
type Action func(eventCtx EventContext) EventType
|
||||
type Action func(ctx context.Context, eventCtx EventContext) EventType
|
||||
|
||||
// Transitions represents a mapping of events and states.
|
||||
type Transitions map[EventType]StateType
|
||||
|
|
@ -95,11 +96,11 @@ type StateMachine struct {
|
|||
|
||||
// ActionEntryFunc is a function that is called before an action is
|
||||
// executed.
|
||||
ActionEntryFunc func(Notification)
|
||||
ActionEntryFunc func(context.Context, Notification)
|
||||
|
||||
// ActionExitFunc is a function that is called after an action is
|
||||
// executed, it is called with the EventType returned by the action.
|
||||
ActionExitFunc func(NextEvent EventType)
|
||||
ActionExitFunc func(ctx context.Context, NextEvent EventType)
|
||||
|
||||
// LastActionError is an error set by the last action executed.
|
||||
LastActionError error
|
||||
|
|
@ -200,7 +201,9 @@ func (s *StateMachine) getNextState(event EventType) (State, error) {
|
|||
// SendEvent sends an event to the state machine. It returns an error if the
|
||||
// event cannot be processed in the current state. Otherwise, it only returns
|
||||
// nil if the event for the last action is a no-op.
|
||||
func (s *StateMachine) SendEvent(event EventType, eventCtx EventContext) error {
|
||||
func (s *StateMachine) SendEvent(ctx context.Context, event EventType,
|
||||
eventCtx EventContext) error {
|
||||
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
|
|
@ -235,7 +238,7 @@ func (s *StateMachine) SendEvent(event EventType, eventCtx EventContext) error {
|
|||
|
||||
// Execute the state machines ActionEntryFunc.
|
||||
if s.ActionEntryFunc != nil {
|
||||
s.ActionEntryFunc(notification)
|
||||
s.ActionEntryFunc(ctx, notification)
|
||||
}
|
||||
|
||||
// Execute the current state's entry function
|
||||
|
|
@ -245,7 +248,7 @@ func (s *StateMachine) SendEvent(event EventType, eventCtx EventContext) error {
|
|||
|
||||
// Execute the next state's action and loop over again if the
|
||||
// event returned is not a no-op.
|
||||
nextEvent := state.Action(eventCtx)
|
||||
nextEvent := state.Action(ctx, eventCtx)
|
||||
|
||||
// Execute the current state's exit function
|
||||
if state.ExitFunc != nil {
|
||||
|
|
@ -254,7 +257,7 @@ func (s *StateMachine) SendEvent(event EventType, eventCtx EventContext) error {
|
|||
|
||||
// Execute the state machines ActionExitFunc.
|
||||
if s.ActionExitFunc != nil {
|
||||
s.ActionExitFunc(nextEvent)
|
||||
s.ActionExitFunc(ctx, nextEvent)
|
||||
}
|
||||
|
||||
// If the next event is a no-op, we're done.
|
||||
|
|
@ -304,7 +307,7 @@ func (s *StateMachine) HandleError(err error) EventType {
|
|||
|
||||
// NoOpAction is a no-op action that can be used by states that don't need to
|
||||
// execute any action.
|
||||
func NoOpAction(_ EventContext) EventType {
|
||||
func NoOpAction(_ context.Context, _ EventContext) EventType {
|
||||
return NoOp
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package fsm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
|
|
@ -22,7 +23,7 @@ type TestStateMachineContext struct {
|
|||
func (c *TestStateMachineContext) GetStates() States {
|
||||
return States{
|
||||
"State1": State{
|
||||
Action: func(ctx EventContext) EventType {
|
||||
Action: func(_ context.Context, ctx EventContext) EventType {
|
||||
return "Event1"
|
||||
},
|
||||
Transitions: Transitions{
|
||||
|
|
@ -30,7 +31,7 @@ func (c *TestStateMachineContext) GetStates() States {
|
|||
},
|
||||
},
|
||||
"State2": State{
|
||||
Action: func(ctx EventContext) EventType {
|
||||
Action: func(_ context.Context, ctx EventContext) EventType {
|
||||
return "NoOp"
|
||||
},
|
||||
Transitions: Transitions{},
|
||||
|
|
@ -39,7 +40,9 @@ func (c *TestStateMachineContext) GetStates() States {
|
|||
}
|
||||
|
||||
// errorAction returns an error.
|
||||
func (c *TestStateMachineContext) errorAction(eventCtx EventContext) EventType {
|
||||
func (c *TestStateMachineContext) errorAction(ctx context.Context,
|
||||
eventCtx EventContext) EventType {
|
||||
|
||||
return c.StateMachine.HandleError(errAction)
|
||||
}
|
||||
|
||||
|
|
@ -58,9 +61,9 @@ func setupTestStateMachineContext() *TestStateMachineContext {
|
|||
// TestStateMachine_Success tests the state machine with a successful event.
|
||||
func TestStateMachine_Success(t *testing.T) {
|
||||
ctx := setupTestStateMachineContext()
|
||||
|
||||
ctxb := context.Background()
|
||||
// Send an event to the state machine.
|
||||
err := ctx.SendEvent("Event1", nil)
|
||||
err := ctx.SendEvent(ctxb, "Event1", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Check that the state machine has transitioned to the next state.
|
||||
|
|
@ -72,8 +75,9 @@ func TestStateMachine_Success(t *testing.T) {
|
|||
func TestStateMachine_ConfigurationError(t *testing.T) {
|
||||
ctx := setupTestStateMachineContext()
|
||||
ctx.StateMachine.States = nil
|
||||
ctxb := context.Background()
|
||||
|
||||
err := ctx.SendEvent("Event1", nil)
|
||||
err := ctx.SendEvent(ctxb, "Event1", nil)
|
||||
require.EqualError(
|
||||
t, err,
|
||||
NewErrConfigError("state machine config is nil").Error(),
|
||||
|
|
@ -83,6 +87,7 @@ func TestStateMachine_ConfigurationError(t *testing.T) {
|
|||
// TestStateMachine_ActionError tests the state machine with an action error.
|
||||
func TestStateMachine_ActionError(t *testing.T) {
|
||||
ctx := setupTestStateMachineContext()
|
||||
ctxb := context.Background()
|
||||
|
||||
states := ctx.StateMachine.States
|
||||
|
||||
|
|
@ -99,13 +104,13 @@ func TestStateMachine_ActionError(t *testing.T) {
|
|||
}
|
||||
|
||||
states["ErrorState"] = State{
|
||||
Action: func(ctx EventContext) EventType {
|
||||
Action: func(_ context.Context, ctx EventContext) EventType {
|
||||
return "NoOp"
|
||||
},
|
||||
Transitions: Transitions{},
|
||||
}
|
||||
|
||||
err := ctx.SendEvent("Event1", nil)
|
||||
err := ctx.SendEvent(ctxb, "Event1", nil)
|
||||
|
||||
// Sending an event to the state machine should not return an error.
|
||||
require.NoError(t, err)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue