From 19d5232293ec131352dc03896d34db82a124f01b Mon Sep 17 00:00:00 2001 From: Janne Valkealahti Date: Sat, 3 Dec 2016 11:26:57 +0000 Subject: [PATCH] Fix fork/join with composite states - Handle new pseudostate listener registration when machine is restored. Make sure that we have only one set of listeners by clearing existing ones. - Fix join issue when regions are partially ended into join source states and machine is persisted from that moment. - Lot of new tests to cover possible cases where these bugs were hidden. - Fixes #286 --- .../listener/OrderedComposite.java | 6 +- .../state/AbstractPseudoState.java | 11 +- .../statemachine/state/ChoicePseudoState.java | 4 + .../statemachine/state/EntryPseudoState.java | 5 + .../statemachine/state/ExitPseudoState.java | 6 + .../statemachine/state/JoinPseudoState.java | 21 + .../state/JunctionPseudoState.java | 4 + .../statemachine/state/PseudoState.java | 11 +- .../support/AbstractStateMachine.java | 13 +- .../persist/StateMachinePersistTests4.java | 572 ++++++++++++++++++ 10 files changed, 647 insertions(+), 6 deletions(-) diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/listener/OrderedComposite.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/listener/OrderedComposite.java index cec69c9d..2e7bfa37 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/listener/OrderedComposite.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/listener/OrderedComposite.java @@ -54,8 +54,10 @@ public class OrderedComposite { public void setItems(List items) { unordered.clear(); ordered.clear(); - for (S s : items) { - add(s); + if (items != null) { + for (S s : items) { + add(s); + } } } diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractPseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractPseudoState.java index 514ed306..ba2027fb 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractPseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractPseudoState.java @@ -15,6 +15,8 @@ */ package org.springframework.statemachine.state; +import java.util.List; + import org.springframework.statemachine.StateContext; /** @@ -49,7 +51,7 @@ public abstract class AbstractPseudoState implements PseudoState { public State entry(StateContext context) { return null; } - + @Override public void exit(StateContext context) { } @@ -58,7 +60,12 @@ public abstract class AbstractPseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { pseudoStateListener.register(listener); } - + + @Override + public void setPseudoStateListeners(List> listeners) { + pseudoStateListener.setListeners(listeners); + } + /** * Notify all {@link PseudoStateListener}s of a new context. * diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ChoicePseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ChoicePseudoState.java index 28a42d6d..706b53d7 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ChoicePseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ChoicePseudoState.java @@ -70,6 +70,10 @@ public class ChoicePseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { } + @Override + public void setPseudoStateListeners(List> listeners) { + } + private boolean evaluateInternal(Guard guard, StateContext context) { try { return guard.evaluate(context); diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/EntryPseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/EntryPseudoState.java index 2be256a1..917841f0 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/EntryPseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/EntryPseudoState.java @@ -15,6 +15,8 @@ */ package org.springframework.statemachine.state; +import java.util.List; + import org.springframework.statemachine.StateContext; /** @@ -56,4 +58,7 @@ public class EntryPseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { } + @Override + public void setPseudoStateListeners(List> listeners) { + } } diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ExitPseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ExitPseudoState.java index 6c76e6d8..a8f00ec4 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ExitPseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/ExitPseudoState.java @@ -15,6 +15,8 @@ */ package org.springframework.statemachine.state; +import java.util.List; + import org.springframework.statemachine.StateContext; import org.springframework.util.Assert; @@ -58,6 +60,10 @@ public class ExitPseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { } + @Override + public void setPseudoStateListeners(List> listeners) { + } + /** * Gets the state holder. * 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 227a55cc..db384c20 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 @@ -16,6 +16,7 @@ package org.springframework.statemachine.state; import java.util.ArrayList; +import java.util.Collection; import java.util.List; import org.apache.commons.logging.Log; @@ -83,6 +84,16 @@ public class JoinPseudoState extends AbstractPseudoState { return joins; } + /** + * Resets join state according to given state ids + * so that we can continue with correct tracking. + * + * @param ids the state id's + */ + public void reset(Collection ids) { + tracker.reset(ids); + } + private boolean evaluateInternal(Guard guard, StateContext context) { try { return guard.evaluate(context); @@ -137,6 +148,16 @@ public class JoinPseudoState extends AbstractPseudoState { notified = false; } + void reset(Collection ids) { + track.clear(); + for (State j : joins) { + if (!ids.contains(j.getId())) { + track.add(j); + } + } + notified = false; + } + public boolean isNotified() { return notified; } diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JunctionPseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JunctionPseudoState.java index ca092f9d..9698d0c9 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JunctionPseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/JunctionPseudoState.java @@ -70,6 +70,10 @@ public class JunctionPseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { } + @Override + public void setPseudoStateListeners(List> listeners) { + } + private boolean evaluateInternal(Guard guard, StateContext context) { try { return guard.evaluate(context); diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/PseudoState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/PseudoState.java index 312f91cd..85633818 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/PseudoState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/PseudoState.java @@ -15,6 +15,8 @@ */ package org.springframework.statemachine.state; +import java.util.List; + import org.springframework.statemachine.StateContext; /** @@ -57,7 +59,7 @@ public interface PseudoState { * @param context the context */ void exit(StateContext context); - + /** * Registers a new {@link PseudoStateListener}. * @@ -65,4 +67,11 @@ public interface PseudoState { */ void addPseudoStateListener(PseudoStateListener listener); + /** + * Registers a new {@link PseudoStateListener}s. Clears all + * existing listeners. + * + * @param listeners the listeners + */ + void setPseudoStateListeners(List> listeners); } 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 5ed5c3ef..4ff55bd6 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 @@ -40,6 +40,7 @@ import org.springframework.statemachine.region.Region; import org.springframework.statemachine.state.AbstractState; import org.springframework.statemachine.state.ForkPseudoState; import org.springframework.statemachine.state.HistoryPseudoState; +import org.springframework.statemachine.state.JoinPseudoState; import org.springframework.statemachine.state.PseudoState; import org.springframework.statemachine.state.PseudoStateContext; import org.springframework.statemachine.state.PseudoStateKind; @@ -355,6 +356,7 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo if (log.isDebugEnabled()) { log.debug("State already set, disabling initial"); } + registerPseudoStateListener(); stateMachineExecutor.setInitialEnabled(false); stateMachineExecutor.start(); // assume that state was set/reseted so we need to @@ -655,6 +657,12 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo } } for (State s : getStates()) { + if (StateMachineUtils.isPseudoState(s, PseudoStateKind.JOIN)) { + JoinPseudoState jps = (JoinPseudoState) s.getPseudoState(); + Collection ids = currentState.getIds(); + jps.reset(ids); + } + // setting history for 'submachines' if (s.isSubmachineState()) { StateMachine submachine = ((AbstractState) s).getSubmachine(); @@ -821,7 +829,8 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo for (State state : states) { PseudoState p = state.getPseudoState(); if (p != null) { - p.addPseudoStateListener(new PseudoStateListener() { + List> listeners = new ArrayList>(); + listeners.add(new PseudoStateListener() { @Override public void onContext(PseudoStateContext context) { PseudoState pseudoState = context.getPseudoState(); @@ -835,6 +844,8 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo pseudoState.exit(stateContext); } }); + // setting instead adding makes sure existing listeners are removed + p.setPseudoStateListeners(listeners); } } } diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/persist/StateMachinePersistTests4.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/persist/StateMachinePersistTests4.java index 229fa2d8..43d345a2 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/persist/StateMachinePersistTests4.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/persist/StateMachinePersistTests4.java @@ -15,13 +15,19 @@ */ package org.springframework.statemachine.persist; +import static org.hamcrest.Matchers.contains; import static org.hamcrest.Matchers.containsInAnyOrder; import static org.hamcrest.Matchers.is; import static org.hamcrest.Matchers.notNullValue; import static org.hamcrest.Matchers.nullValue; import static org.junit.Assert.assertThat; +import java.util.ArrayList; +import java.util.Collection; import java.util.HashMap; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; import org.junit.Test; import org.springframework.context.annotation.AnnotationConfigApplicationContext; @@ -30,11 +36,18 @@ import org.springframework.statemachine.AbstractStateMachineTests; import org.springframework.statemachine.StateMachine; import org.springframework.statemachine.StateMachineContext; import org.springframework.statemachine.StateMachinePersist; +import org.springframework.statemachine.TestUtils; import org.springframework.statemachine.config.EnableStateMachineFactory; import org.springframework.statemachine.config.EnumStateMachineConfigurerAdapter; import org.springframework.statemachine.config.StateMachineFactory; import org.springframework.statemachine.config.builders.StateMachineStateConfigurer; import org.springframework.statemachine.config.builders.StateMachineTransitionConfigurer; +import org.springframework.statemachine.listener.OrderedComposite; +import org.springframework.statemachine.listener.StateMachineListenerAdapter; +import org.springframework.statemachine.state.CompositePseudoStateListener; +import org.springframework.statemachine.state.PseudoState; +import org.springframework.statemachine.state.State; +import org.springframework.statemachine.transition.Transition; public class StateMachinePersistTests4 extends AbstractStateMachineTests { @@ -93,6 +106,409 @@ public class StateMachinePersistTests4 extends AbstractStateMachineTests { assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); } + @Test + public void testJoinAfterPersistRegionsNotEnteredJoinStates() throws Exception { + context.register(Config1.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = stateMachineFactory.getStateMachine(); + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsNotEnteredJoinStatesRestoreTwice() throws Exception { + context.register(Config1.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = stateMachineFactory.getStateMachine(); + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsPartialEnteredJoinStates() throws Exception { + context.register(Config1.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsPartialEnteredJoinStatesRestoreTwice() throws Exception { + context.register(Config1.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsNotEnteredJoinStatesWithEnds() throws Exception { + context.register(Config2.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = stateMachineFactory.getStateMachine(); + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsNotEnteredJoinStatesRestoreTwiceWithEnds() throws Exception { + context.register(Config2.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = stateMachineFactory.getStateMachine(); + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsPartialEnteredJoinStatesWithEnds() throws Exception { + context.register(Config2.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + @Test + public void testJoinAfterPersistRegionsPartialEnteredJoinStatesRestoreTwiceWithEnds() throws Exception { + context.register(Config2.class); + context.refresh(); + + @SuppressWarnings("unchecked") + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + StateMachine stateMachine = stateMachineFactory.getStateMachine("testid"); + stateMachine.start(); + + stateMachine.sendEvent(TestEvents.E1); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E2); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + persister.persist(stateMachine, "xxx1"); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + + stateMachine = persister.restore(stateMachine, "xxx1"); + assertThat(stateMachine.getId(), is("testid")); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + stateMachine.sendEvent(TestEvents.E3); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder(TestStates.S4)); + } + + + @Test + @SuppressWarnings("unchecked") + public void testJoinFromSuperAfterPersistRegions() throws Exception { + context.register(Config3.class); + context.refresh(); + + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + StateMachine machine = stateMachineFactory.getStateMachine("testid"); + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + 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)); + assertThat(machine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + persister.persist(machine, "xxx1"); + machine = persister.restore(machine, "xxx1"); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + assertThat(machine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + listener.reset(2); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(2)); + + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + + // try fresh machine + machine = stateMachineFactory.getStateMachine("testid"); + machine = persister.restore(machine, "xxx1"); + machine.addStateListener(listener); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + assertThat(machine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + listener.reset(2); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(2)); + + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + } + + @Test + @SuppressWarnings("unchecked") + public void testJoinFromSuperAfterPersistRegionsPartial() throws Exception { + context.register(Config3.class); + context.refresh(); + + StateMachineFactory stateMachineFactory = context.getBean(StateMachineFactory.class); + StateMachine machine = stateMachineFactory.getStateMachine("testid"); + InMemoryStateMachinePersist stateMachinePersist = new InMemoryStateMachinePersist(); + StateMachinePersister persister = new DefaultStateMachinePersister<>(stateMachinePersist); + + TestListener listener = new TestListener(); + machine.addStateListener(listener); + listener.reset(1); + assertThat(machine, notNullValue()); + machine.start(); + assertPseudoStatesHaveOneListener(machine); + 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)); + assertThat(machine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S20, TestStates.S30)); + + listener.reset(1); + machine.sendEvent(TestEvents.E2); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(1)); + assertThat(machine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + + persister.persist(machine, "xxx1"); + + machine = persister.restore(machine, "xxx1"); + assertPseudoStatesHaveOneListener(machine); + listener.reset(2); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(2)); + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + + + machine = persister.restore(machine, "xxx1"); + assertPseudoStatesHaveOneListener(machine); + listener.reset(2); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(2)); + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + + // try fresh machine + machine = stateMachineFactory.getStateMachine("testid"); + machine = persister.restore(machine, "xxx1"); + machine.addStateListener(listener); + assertThat(machine.getState().getIds(), containsInAnyOrder(TestStates.S2, TestStates.S21, TestStates.S30)); + listener.reset(2); + machine.sendEvent(TestEvents.E3); + assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); + assertThat(listener.stateChangedCount, is(2)); + assertThat(machine.getState().getIds(), contains(TestStates.S4)); + assertPseudoStatesHaveOneListener(machine); + } + + private void assertPseudoStatesHaveOneListener(Object machine) throws Exception { + Collection> states = TestUtils.readField("states", machine); + for (State s : states) { + PseudoState ps = s.getPseudoState(); + if (ps != null) { + CompositePseudoStateListener pseudoStateListener = TestUtils.readField("pseudoStateListener", ps); + OrderedComposite listeners = TestUtils.readField("listeners", pseudoStateListener); + List list = TestUtils.readField("list", listeners); + assertThat(list.size(), is(1)); + } + } + } + @Configuration @EnableStateMachineFactory static class Config1 extends EnumStateMachineConfigurerAdapter { @@ -157,6 +573,162 @@ public class StateMachinePersistTests4 extends AbstractStateMachineTests { } + @Configuration + @EnableStateMachineFactory + static class Config2 extends EnumStateMachineConfigurerAdapter { + + @Override + public void configure(StateMachineStateConfigurer states) throws Exception { + states + .withStates() + .initial(TestStates.SI) + .state(TestStates.SI) + .fork(TestStates.S1) + .state(TestStates.S2) + .end(TestStates.SF) + .join(TestStates.S3) + .end(TestStates.S4) + .and() + .withStates() + .parent(TestStates.S2) + .initial(TestStates.S20) + .state(TestStates.S20) + .end(TestStates.S21) + .and() + .withStates() + .parent(TestStates.S2) + .initial(TestStates.S30) + .state(TestStates.S30) + .end(TestStates.S31); + } + + @Override + public void configure(StateMachineTransitionConfigurer transitions) throws Exception { + transitions + .withExternal() + .source(TestStates.SI) + .target(TestStates.S2) + .event(TestEvents.E1) + .and() + .withExternal() + .source(TestStates.S20) + .target(TestStates.S21) + .event(TestEvents.E2) + .and() + .withExternal() + .source(TestStates.S30) + .target(TestStates.S31) + .event(TestEvents.E3) + .and() + .withFork() + .source(TestStates.S1) + .target(TestStates.S20) + .target(TestStates.S30) + .and() + .withJoin() + .source(TestStates.S21) + .source(TestStates.S31) + .target(TestStates.S3) + .and() + .withExternal() + .source(TestStates.S3) + .target(TestStates.S4); + } + + } + + @Configuration + @EnableStateMachineFactory + static class Config3 extends EnumStateMachineConfigurerAdapter { + + @Override + public void configure(StateMachineStateConfigurer states) throws Exception { + states + .withStates() + .initial(TestStates.SI) + .state(TestStates.S2) + .join(TestStates.S3) + .state(TestStates.S4) + .and() + .withStates() + .parent(TestStates.S2) + .initial(TestStates.S20) + .end(TestStates.S21) + .and() + .withStates() + .parent(TestStates.S2) + .initial(TestStates.S30) + .end(TestStates.S31); + } + + @Override + public void configure(StateMachineTransitionConfigurer transitions) throws Exception { + transitions + .withExternal() + .source(TestStates.SI) + .target(TestStates.S2) + .event(TestEvents.E1) + .and() + .withExternal() + .source(TestStates.S20) + .target(TestStates.S21) + .event(TestEvents.E2) + .and() + .withExternal() + .source(TestStates.S30) + .target(TestStates.S31) + .event(TestEvents.E3) + .and() + .withJoin() + .source(TestStates.S2) + .target(TestStates.S3) + .and() + .withExternal() + .source(TestStates.S3) + .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>(); + final List> tos = new ArrayList<>(); + + @Override + public void stateChanged(State from, State to) { + tos.add(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(); + tos.clear(); + } + } + static class InMemoryStateMachinePersist implements StateMachinePersist { private final HashMap> contexts = new HashMap<>();