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 38b52801..12bb7a2b 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 @@ -494,6 +494,9 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo @Override public void resetStateMachine(StateMachineContext stateMachineContext) { if (stateMachineContext == null) { + log.info("Got null context, resetting to initial state and clearing extended state"); + currentState = initialState; + extendedState.getVariables().clear(); return; } if (log.isDebugEnabled()) { 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 553b4268..d874cf2f 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 @@ -223,6 +223,34 @@ public class StateMachineResetTests extends AbstractStateMachineTests { assertThat((Integer)machine.getExtendedState().getVariables().get("count"), is(2)); } + @Test + public void testResetWithNullContext() 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)); + assertThat((Integer)machine.getExtendedState().getVariables().get("foo"), is(0)); + + machine.sendEvent(Events.I); + assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S1, States.S12)); + assertThat((Integer)machine.getExtendedState().getVariables().get("foo"), is(0)); + + machine.stop(); + machine.getStateMachineAccessor().doWithAllRegions(new StateMachineFunction>() { + + @Override + public void apply(StateMachineAccess function) { + function.resetStateMachine(null); + } + }); + machine.start(); + assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S1, States.S11)); + assertThat(machine.getExtendedState().getVariables().size(), is(0)); + } + @Configuration @EnableStateMachine static class Config1 extends EnumStateMachineConfigurerAdapter {