diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java index 60ee5491..b500fbfe 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java @@ -296,7 +296,7 @@ public class EnumStateMachineFactory, E extends Enum> exten if (state != null) { states.add(state); if (stateData.isInitial()) { - initialState = state; + initialState = state; } continue; } @@ -373,8 +373,27 @@ public class EnumStateMachineFactory, E extends Enum> exten S s = stateData.getState(); List list = stateMachineTransitions.getJoins().get(s); List> joins = new ArrayList>(); - for (S fs : list) { - joins.add(stateMap.get(fs)); + + // if join source is a regionstate, get + // it's end states from regions + if (list.size() == 1) { + State ss1 = stateMap.get(list.get(0)); + if (ss1 instanceof RegionState) { + Collection> regions = ((RegionState)ss1).getRegions(); + for (Region r : regions) { + Collection> ss2 = r.getStates(); + for (State ss3 : ss2) { + if (ss3.getPseudoState() != null && ss3.getPseudoState().getKind() == PseudoStateKind.END) { + joins.add(ss3); + continue; + } + } + } + } + } else { + for (S fs : list) { + joins.add(stateMap.get(fs)); + } } JoinPseudoState pseudoState = new JoinPseudoState(joins); state = new EnumState(stateData.getState(), stateData.getDeferred(), diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JoinPseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JoinPseudoState.java index 7e176ad7..052dffa5 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JoinPseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JoinPseudoState.java @@ -15,6 +15,7 @@ */ package org.springframework.statemachine.state; +import java.util.ArrayList; import java.util.List; import org.springframework.statemachine.StateContext; @@ -42,13 +43,16 @@ public class JoinPseudoState extends AbstractPseudoState { @Override public State entry(StateContext context) { - tracker = new JoinTracker(this, joins); + tracker = new JoinTracker(this, new ArrayList>(joins)); context.getStateMachine().addStateListener(tracker); return null; } - + @Override public void exit(StateContext context) { + if (context != null) { + context.getStateMachine().removeStateListener(tracker); + } tracker = null; } @@ -60,9 +64,8 @@ public class JoinPseudoState extends AbstractPseudoState { private final PseudoState pseudoState; private final List> track; - // TOOO use flat till we can unregister listener - private boolean done = false; - + private boolean notified = false; + public JoinTracker(PseudoState pseudoState, List> track) { this.pseudoState = pseudoState; this.track = track; @@ -70,13 +73,12 @@ public class JoinPseudoState extends AbstractPseudoState { @Override public void stateChanged(State from, State to) { - if (done) { - return; - } - track.remove(to); - if (track.size() == 0) { - done = true; - notifyContext(new DefaultPseudoStateContext(pseudoState, PseudoAction.JOIN_COMPLETED)); + if (!notified && track.size() > 0) { + track.remove(to); + if (track.size() == 0) { + notified = true; + notifyContext(new DefaultPseudoStateContext(pseudoState, PseudoAction.JOIN_COMPLETED)); + } } } 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 ef693597..c071bb43 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 @@ -370,7 +370,8 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo public void onContext(PseudoStateContext context) { PseudoState pseudoState = context.getPseudoState(); State toState = findStateWithPseudoState(pseudoState); - pseudoState.exit(null); + StateContext stateContext = buildStateContext(null, null, AbstractStateMachine.this); + pseudoState.exit(stateContext); switchToState(toState, null, null, AbstractStateMachine.this); } }); diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/JoinStateTests.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/JoinStateTests.java index ef2c04f5..5c784162 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/JoinStateTests.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/state/JoinStateTests.java @@ -16,9 +16,15 @@ package org.springframework.statemachine.state; import static org.hamcrest.Matchers.contains; +import static org.hamcrest.Matchers.is; import static org.hamcrest.Matchers.notNullValue; import static org.junit.Assert.assertThat; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; + import org.junit.Test; import org.springframework.context.annotation.AnnotationConfigApplicationContext; import org.springframework.context.annotation.Configuration; @@ -29,6 +35,8 @@ import org.springframework.statemachine.config.EnableStateMachine; import org.springframework.statemachine.config.EnumStateMachineConfigurerAdapter; import org.springframework.statemachine.config.builders.StateMachineStateConfigurer; import org.springframework.statemachine.config.builders.StateMachineTransitionConfigurer; +import org.springframework.statemachine.listener.StateMachineListenerAdapter; +import org.springframework.statemachine.transition.Transition; public class JoinStateTests extends AbstractStateMachineTests { @@ -39,11 +47,46 @@ public class JoinStateTests extends AbstractStateMachineTests { @Test @SuppressWarnings("unchecked") - public void testJoin() { + public void testJoin() throws Exception { context.register(BaseConfig.class, Config1.class); context.refresh(); EnumStateMachine machine = context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, EnumStateMachine.class); + TestListener listener = new TestListener(); + machine.addStateListener(listener); + listener.reset(1); + assertThat(machine, notNullValue()); + machine.start(); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + + listener.reset(3); + machine.sendEvent(TestEvents.E1); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + + listener.reset(3); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + } + + @Test + @SuppressWarnings("unchecked") + public void testJoinLoopTwice() throws Exception { + context.register(BaseConfig.class, Config1.class); + context.refresh(); + EnumStateMachine machine = + context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, EnumStateMachine.class); + TestListener listener = new TestListener(); + machine.addStateListener(listener); assertThat(machine, notNullValue()); machine.start(); machine.sendEvent(TestEvents.E1); @@ -51,15 +94,73 @@ public class JoinStateTests extends AbstractStateMachineTests { machine.sendEvent(TestEvents.E3); assertThat(machine.getState().getIds(), contains(TestStates.S4)); + + listener.reset(1); + machine.sendEvent(TestEvents.E4); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + assertThat(machine.getState().getIds(), contains(TestStates.SI)); + + listener.reset(3); + machine.sendEvent(TestEvents.E1); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + + listener.reset(3); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + assertThat(machine.getState().getIds(), contains(TestStates.S4)); } @Test @SuppressWarnings("unchecked") - public void testJoinSuper() { + public void testJoinSuper() throws Exception { context.register(BaseConfig.class, Config2.class); context.refresh(); EnumStateMachine machine = context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, EnumStateMachine.class); + TestListener listener = new TestListener(); + machine.addStateListener(listener); + listener.reset(1); + assertThat(machine, notNullValue()); + machine.start(); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + + listener.reset(3); + machine.sendEvent(TestEvents.E1); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + + listener.reset(3); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + } + + @Test + @SuppressWarnings("unchecked") + public void testJoinSuperLoopTwice() throws Exception { + context.register(BaseConfig.class, Config2.class); + context.refresh(); + EnumStateMachine machine = + context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, EnumStateMachine.class); + TestListener listener = new TestListener(); + machine.addStateListener(listener); assertThat(machine, notNullValue()); machine.start(); machine.sendEvent(TestEvents.E1); @@ -67,6 +168,29 @@ public class JoinStateTests extends AbstractStateMachineTests { machine.sendEvent(TestEvents.E3); assertThat(machine.getState().getIds(), contains(TestStates.S4)); + + listener.reset(1); + machine.sendEvent(TestEvents.E4); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + assertThat(machine.getState().getIds(), contains(TestStates.SI)); + + listener.reset(3); + machine.sendEvent(TestEvents.E1); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + + listener.reset(3); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(3)); + + assertThat(machine.getState().getIds(), contains(TestStates.S4)); } @Configuration @@ -78,9 +202,7 @@ public class JoinStateTests extends AbstractStateMachineTests { states .withStates() .initial(TestStates.SI) - .state(TestStates.SI) .state(TestStates.S2) - .end(TestStates.SF) .join(TestStates.S3) .state(TestStates.S4) .and() @@ -122,7 +244,12 @@ public class JoinStateTests extends AbstractStateMachineTests { .and() .withExternal() .source(TestStates.S3) - .target(TestStates.S4); + .target(TestStates.S4) + .and() + .withExternal() + .source(TestStates.S4) + .target(TestStates.SI) + .event(TestEvents.E4); } } @@ -136,23 +263,19 @@ public class JoinStateTests extends AbstractStateMachineTests { states .withStates() .initial(TestStates.SI) - .state(TestStates.SI) .state(TestStates.S2) - .end(TestStates.SF) .join(TestStates.S3) .state(TestStates.S4) .and() .withStates() .parent(TestStates.S2) .initial(TestStates.S20) - .state(TestStates.S20) - .state(TestStates.S21) + .end(TestStates.S21) .and() .withStates() .parent(TestStates.S2) .initial(TestStates.S30) - .state(TestStates.S30) - .state(TestStates.S31); + .end(TestStates.S31); } @Override @@ -179,7 +302,44 @@ public class JoinStateTests extends AbstractStateMachineTests { .and() .withExternal() .source(TestStates.S3) - .target(TestStates.S4); + .target(TestStates.S4) + .and() + .withExternal() + .source(TestStates.S4) + .target(TestStates.SI) + .event(TestEvents.E4); + } + + } + + private static class TestListener extends StateMachineListenerAdapter { + + volatile CountDownLatch stateChangedLatch = new CountDownLatch(1); + volatile CountDownLatch transitionLatch = new CountDownLatch(0); + volatile int stateChangedCount = 0; + final List> transitions = new ArrayList>(); + + @Override + public void stateChanged(State from, State to) { + stateChangedCount++; + stateChangedLatch.countDown(); + } + + @Override + public void transition(Transition transition) { + transitions.add(transition); + transitionLatch.countDown(); + } + + public void reset(int c1) { + reset(c1, 0); + } + + public void reset(int c1, int c2) { + stateChangedLatch = new CountDownLatch(c1); + transitionLatch = new CountDownLatch(c2); + stateChangedCount = 0; + transitions.clear(); } }