diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateMachineState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateMachineState.java index 9834f0f3..19bc500d 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateMachineState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateMachineState.java @@ -216,6 +216,16 @@ public class StateMachineState extends AbstractState { getSubmachine().getStateMachineAccessor().doWithRegion( new StateMachineFunction>() { + @Override + public void apply(StateMachineAccess function) { + function.setInitialEnabled(false); + } + }); + } + if (immediateDeepParent == null && getSubmachine().getStates().contains(target) && isEntry(target)) { + getSubmachine().getStateMachineAccessor().doWithRegion( + new StateMachineFunction>() { + @Override public void apply(StateMachineAccess function) { function.setInitialEnabled(false); @@ -230,6 +240,10 @@ public class StateMachineState extends AbstractState { return state.getPseudoState() != null && state.getPseudoState().getKind() == PseudoStateKind.INITIAL; } + private boolean isEntry(State state) { + return state.getPseudoState() != null && state.getPseudoState().getKind() == PseudoStateKind.ENTRY; + } + private State findDeepParent(Collection> states, State state) { for (State s : states) { if (s.getStates().contains(state)) { diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/ExitEntryStateTests.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/ExitEntryStateTests.java index 6ab44a63..33de1e30 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/ExitEntryStateTests.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/ExitEntryStateTests.java @@ -19,6 +19,9 @@ import static org.hamcrest.Matchers.contains; import static org.hamcrest.Matchers.notNullValue; import static org.junit.Assert.assertThat; +import java.util.ArrayList; +import java.util.List; + import org.junit.Test; import org.springframework.context.annotation.AnnotationConfigApplicationContext; import org.springframework.context.annotation.Configuration; @@ -29,6 +32,7 @@ import org.springframework.statemachine.config.EnableStateMachine; import org.springframework.statemachine.config.StateMachineConfigurerAdapter; import org.springframework.statemachine.config.builders.StateMachineStateConfigurer; import org.springframework.statemachine.config.builders.StateMachineTransitionConfigurer; +import org.springframework.statemachine.listener.StateMachineListenerAdapter; public class ExitEntryStateTests extends AbstractStateMachineTests { @@ -40,14 +44,38 @@ public class ExitEntryStateTests extends AbstractStateMachineTests { StateMachine machine = context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, StateMachine.class); assertThat(machine, notNullValue()); + TestStateEntryExitListener listener = new TestStateEntryExitListener(); + machine.addStateListener(listener); machine.start(); assertThat(machine.getState().getIds(), contains("S1")); + listener.reset(); machine.sendEvent("ENTRY1"); assertThat(machine.getState().getIds(), contains("S2", "S22")); + assertThat(listener.exited, contains("S1")); + assertThat(listener.entered, contains("S2", "S22")); machine.sendEvent("EXIT1"); assertThat(machine.getState().getIds(), contains("S4")); } + @SuppressWarnings("unchecked") + @Test + public void testSimpleEntryToInitial() { + context.register(Config1.class); + context.refresh(); + StateMachine machine = + context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, StateMachine.class); + assertThat(machine, notNullValue()); + TestStateEntryExitListener listener = new TestStateEntryExitListener(); + machine.addStateListener(listener); + machine.start(); + assertThat(machine.getState().getIds(), contains("S1")); + listener.reset(); + machine.sendEvent("ENTRY3"); + assertThat(machine.getState().getIds(), contains("S2", "S21")); + assertThat(listener.exited, contains("S1")); + assertThat(listener.entered, contains("S2", "S21")); + } + @SuppressWarnings("unchecked") @Test public void testMultipleExitsToSameState() { @@ -92,6 +120,7 @@ public class ExitEntryStateTests extends AbstractStateMachineTests { .initial("S21") .entry("S2ENTRY1") .entry("S2ENTRY2") + .entry("S2ENTRY3") .exit("S2EXIT1") .exit("S2EXIT2") .state("S22") @@ -117,6 +146,10 @@ public class ExitEntryStateTests extends AbstractStateMachineTests { .source("S1").target("S2ENTRY2") .event("ENTRY2") .and() + .withExternal() + .source("S1").target("S2ENTRY3") + .event("ENTRY3") + .and() .withExternal() .source("S22").target("S2EXIT1") .event("EXIT1") @@ -131,6 +164,9 @@ public class ExitEntryStateTests extends AbstractStateMachineTests { .withEntry() .source("S2ENTRY2").target("S23") .and() + .withEntry() + .source("S2ENTRY3").target("S21") + .and() .withExit() .source("S2EXIT1").target("S4") .and() @@ -197,6 +233,26 @@ public class ExitEntryStateTests extends AbstractStateMachineTests { } } + private static class TestStateEntryExitListener extends StateMachineListenerAdapter { + + List entered = new ArrayList<>(); + List exited = new ArrayList<>(); + + @Override + public void stateEntered(State state) { + entered.add(state.getId()); + } + + @Override + public void stateExited(State state) { + exited.add(state.getId()); + } + + public void reset() { + entered.clear(); + exited.clear(); + } + } @Override protected AnnotationConfigApplicationContext buildContext() {