diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/AbstractStateMachine.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/AbstractStateMachine.java index 81122240..6f3e0590 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/AbstractStateMachine.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/AbstractStateMachine.java @@ -259,6 +259,7 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo for (final State state : states) { state.addStateListener(new StateListenerAdapter() { + @Override public void onComplete(StateContext context) { ((AbstractStateMachine)getRelayStateMachine()).executeTriggerlessTransitions(AbstractStateMachine.this, context, state); }; @@ -632,6 +633,7 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo return buf.toString(); } + @SuppressWarnings("rawtypes") @Override public void resetStateMachine(StateMachineContext stateMachineContext) { // TODO: this function needs a serious rewrite @@ -651,7 +653,14 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo // handle state reset for (State s : getStates()) { for (State ss : s.getStates()) { - if (state != null && ss.getIds().contains(state)) { + + boolean enumMatch = false; + if (state instanceof Enum && ss.getId() instanceof Enum + && ((Enum) ss.getId()).ordinal() == ((Enum) state).ordinal()) { + enumMatch = true; + } + + if (state != null && (ss.getIds().contains(state) || enumMatch) ) { currentState = s; // setting lastState here is needed for restore lastState = currentState; @@ -706,7 +715,12 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo } else { for (final StateMachineContext child : stateMachineContext.getChilds()) { S state2 = child.getState(); - if (state2 != null && ss.getIds().contains(state2)) { + boolean enumMatch2 = false; + if (state2 instanceof Enum && ss.getId() instanceof Enum + && ((Enum) ss.getId()).ordinal() == ((Enum) state2).ordinal()) { + enumMatch2 = true; + } + if (state2 != null && (ss.getIds().contains(state2) || enumMatch2) ) { currentState = s; lastState = currentState; stateSet = true; diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineResetTests.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineResetTests.java index 6dc55bbf..d1090103 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineResetTests.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineResetTests.java @@ -18,6 +18,8 @@ package org.springframework.statemachine; import static org.hamcrest.Matchers.containsInAnyOrder; import static org.hamcrest.Matchers.is; import static org.hamcrest.Matchers.nullValue; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotEquals; import static org.junit.Assert.assertThat; import java.util.ArrayList; @@ -255,6 +257,37 @@ public class StateMachineResetTests extends AbstractStateMachineTests { assertThat(machine.getExtendedState().getVariables().size(), is(0)); } + @Test + public void testResetWithEnumToCorrectStartState() throws Exception { + context.register(Config1.class); + context.refresh(); + @SuppressWarnings("unchecked") + StateMachine machine = context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, + StateMachine.class); + + machine.start(); + assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S1, States.S11)); + + machine.sendEvent(Events.I); + assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S1, States.S12)); + + machine.stop(); + DefaultStateMachineContext stateMachineContext = new DefaultStateMachineContext( + States.S11, null, null, null); + machine.getStateMachineAccessor() + .doWithAllRegions(new StateMachineFunction>() { + + @Override + public void apply(StateMachineAccess function) { + function.resetStateMachine(stateMachineContext); + } + }); + machine.start(); + assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S1, States.S11)); + assertEquals(States.S11, stateMachineContext.getState()); + assertNotEquals(stateMachineContext.getState(), machine.getInitialState()); + } + @Test public void testRestoreWithTimer() throws Exception { context.register(Config4.class);