From 53d6ddefa9bdb1d7e63a91b1702db502de4feffd Mon Sep 17 00:00:00 2001 From: Janne Valkealahti Date: Mon, 11 Mar 2019 11:17:13 +0000 Subject: [PATCH] Fix StateMachineInterceptor calls - Due to a regression when new preStateChange method were added to pass on root machine, old method did not get called anymore. Now calling that as well and fixing some other things around it. - Same changes for postStateChange. - Fixes #633 --- .../ensemble/DistributedStateMachine.java | 11 +- .../support/StateMachineInterceptorList.java | 17 +-- .../support/StateChangeInterceptorTests.java | 135 ++++++++++++------ .../persist/PersistStateMachineHandler.java | 6 - .../recipes/tasks/TasksHandler.java | 10 +- 5 files changed, 99 insertions(+), 80 deletions(-) diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/ensemble/DistributedStateMachine.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/ensemble/DistributedStateMachine.java index 366acd99..679311f1 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/ensemble/DistributedStateMachine.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/ensemble/DistributedStateMachine.java @@ -199,6 +199,11 @@ public class DistributedStateMachine extends LifecycleObjectSupport implem @Override public void preStateChange(State state, Message message, Transition transition, StateMachine stateMachine) { + } + + @Override + public void preStateChange(State state, Message message, Transition transition, + StateMachine stateMachine, StateMachine rootStateMachine) { if (log.isTraceEnabled()) { log.trace("Received preStateChange from " + stateMachine + " for delegate " + delegate); } @@ -211,12 +216,6 @@ public class DistributedStateMachine extends LifecycleObjectSupport implem } } - @Override - public void preStateChange(State state, Message message, Transition transition, - StateMachine stateMachine, StateMachine rootStateMachine) { - preStateChange(state, message, transition, stateMachine); - } - @Override public void postStateChange(State state, Message message, Transition transition, StateMachine stateMachine) { diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineInterceptorList.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineInterceptorList.java index 0aa84a54..b3c76ba3 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineInterceptorList.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineInterceptorList.java @@ -94,17 +94,12 @@ public class StateMachineInterceptorList { * @param message the message * @param transition the transition * @param stateMachine the state machine + * @param rootStateMachine the root state machine */ - public void preStateChange(State state, Message message, Transition transition, - StateMachine stateMachine) { - for (StateMachineInterceptor interceptor : interceptors) { - interceptor.preStateChange(state, message, transition, stateMachine); - } - } - public void preStateChange(State state, Message message, Transition transition, StateMachine stateMachine, StateMachine rootStateMachine) { for (StateMachineInterceptor interceptor : interceptors) { + interceptor.preStateChange(state, message, transition, stateMachine); interceptor.preStateChange(state, message, transition, stateMachine, rootStateMachine); } } @@ -116,17 +111,13 @@ public class StateMachineInterceptorList { * @param message the message * @param transition the transition * @param stateMachine the state machine + * @param rootStateMachine the root state machine */ - public void postStateChange(State state, Message message, Transition transition, - StateMachine stateMachine) { - for (StateMachineInterceptor interceptor : interceptors) { - interceptor.postStateChange(state, message, transition, stateMachine); - } - } public void postStateChange(State state, Message message, Transition transition, StateMachine stateMachine, StateMachine rootStateMachine) { for (StateMachineInterceptor interceptor : interceptors) { + interceptor.postStateChange(state, message, transition, stateMachine); interceptor.postStateChange(state, message, transition, stateMachine, rootStateMachine); } } diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/support/StateChangeInterceptorTests.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/support/StateChangeInterceptorTests.java index fc9ac8d5..ddb0f7d3 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/support/StateChangeInterceptorTests.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/support/StateChangeInterceptorTests.java @@ -83,8 +83,10 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { machine.sendEvent(Events.C); assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(3)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S2, States.S21, States.S211)); machine.sendEvent(Events.H); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S0, States.S2, States.S21, States.S211)); @@ -120,8 +122,10 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S1)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); interceptor.reset(1); listener.reset(1); @@ -129,8 +133,10 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S2)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); } @Test @@ -162,8 +168,10 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S2)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); } @Test @@ -195,13 +203,20 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S2)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); - assertThat(interceptor.postStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.postStateChangeCount, is(1)); - assertThat(interceptor.preStateChangeStates.size(), is(1)); - assertThat(interceptor.postStateChangeStates.size(), is(1)); - assertThat(interceptor.preStateChangeStates.get(0).getId(), is(interceptor.postStateChangeStates.get(0).getId())); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.postStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.postStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeStates1.size(), is(1)); + assertThat(interceptor.postStateChangeStates1.size(), is(1)); + assertThat(interceptor.preStateChangeStates1.get(0).getId(), is(interceptor.postStateChangeStates1.get(0).getId())); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); + assertThat(interceptor.postStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.postStateChangeCount2, is(1)); + assertThat(interceptor.preStateChangeStates2.size(), is(1)); + assertThat(interceptor.postStateChangeStates2.size(), is(1)); + assertThat(interceptor.preStateChangeStates2.get(0).getId(), is(interceptor.postStateChangeStates2.get(0).getId())); } @Test @@ -233,13 +248,20 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S3)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); - assertThat(interceptor.postStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.postStateChangeCount, is(1)); - assertThat(interceptor.preStateChangeStates.size(), is(1)); - assertThat(interceptor.postStateChangeStates.size(), is(1)); - assertThat(interceptor.preStateChangeStates.get(0).getId(), is(interceptor.postStateChangeStates.get(0).getId())); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.postStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.postStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeStates1.size(), is(1)); + assertThat(interceptor.postStateChangeStates1.size(), is(1)); + assertThat(interceptor.preStateChangeStates1.get(0).getId(), is(interceptor.postStateChangeStates1.get(0).getId())); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); + assertThat(interceptor.postStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.postStateChangeCount2, is(1)); + assertThat(interceptor.preStateChangeStates2.size(), is(1)); + assertThat(interceptor.postStateChangeStates2.size(), is(1)); + assertThat(interceptor.preStateChangeStates2.get(0).getId(), is(interceptor.postStateChangeStates2.get(0).getId())); } @Test @@ -271,8 +293,10 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S1)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); interceptor.reset(1); listener.reset(1); @@ -280,8 +304,10 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); assertThat(machine.getState().getIds(), containsInAnyOrder(States.S2)); - assertThat(interceptor.preStateChangeLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(interceptor.preStateChangeCount, is(1)); + assertThat(interceptor.preStateChangeLatch1.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount1, is(1)); + assertThat(interceptor.preStateChangeLatch2.await(2, TimeUnit.SECONDS), is(true)); + assertThat(interceptor.preStateChangeCount2, is(1)); } @Configuration @@ -598,12 +624,18 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { private static class TestStateChangeInterceptor implements StateMachineInterceptor { - volatile CountDownLatch preStateChangeLatch = new CountDownLatch(1); - volatile CountDownLatch postStateChangeLatch = new CountDownLatch(1); - volatile int preStateChangeCount = 0; - volatile int postStateChangeCount = 0; - ArrayList> preStateChangeStates = new ArrayList<>(); - ArrayList> postStateChangeStates = new ArrayList<>(); + volatile CountDownLatch preStateChangeLatch1 = new CountDownLatch(1); + volatile CountDownLatch preStateChangeLatch2 = new CountDownLatch(1); + volatile CountDownLatch postStateChangeLatch1 = new CountDownLatch(1); + volatile CountDownLatch postStateChangeLatch2 = new CountDownLatch(1); + volatile int preStateChangeCount1 = 0; + volatile int preStateChangeCount2 = 0; + volatile int postStateChangeCount1 = 0; + volatile int postStateChangeCount2 = 0; + ArrayList> preStateChangeStates1 = new ArrayList<>(); + ArrayList> preStateChangeStates2 = new ArrayList<>(); + ArrayList> postStateChangeStates1 = new ArrayList<>(); + ArrayList> postStateChangeStates2 = new ArrayList<>(); @Override public Message preEvent(Message message, StateMachine stateMachine) { @@ -613,32 +645,35 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { @Override public void preStateChange(State state, Message message, Transition transition, StateMachine stateMachine) { - preStateChangeStates.add(state); - preStateChangeCount++; - preStateChangeLatch.countDown(); - + preStateChangeStates1.add(state); + preStateChangeCount1++; + preStateChangeLatch1.countDown(); } @Override public void preStateChange(State state, Message message, Transition transition, StateMachine stateMachine, StateMachine rootStateMachine) { - preStateChange(state, message, transition, stateMachine); + preStateChangeStates2.add(state); + preStateChangeCount2++; + preStateChangeLatch2.countDown(); } @Override public void postStateChange(State state, Message message, Transition transition, StateMachine stateMachine) { - postStateChangeStates.add(state); - postStateChangeCount++; - postStateChangeLatch.countDown(); + postStateChangeStates1.add(state); + postStateChangeCount1++; + postStateChangeLatch1.countDown(); } @Override public void postStateChange(State state, Message message, Transition transition, StateMachine stateMachine, StateMachine rootStateMachine) { - postStateChange(state, message, transition, stateMachine); + postStateChangeStates2.add(state); + postStateChangeCount2++; + postStateChangeLatch2.countDown(); } @Override @@ -653,12 +688,18 @@ public class StateChangeInterceptorTests extends AbstractStateMachineTests { public void reset(int c1) { - preStateChangeLatch = new CountDownLatch(c1); - preStateChangeCount = 0; - postStateChangeLatch = new CountDownLatch(c1); - postStateChangeCount = 0; - preStateChangeStates.clear(); - postStateChangeStates.clear(); + preStateChangeLatch1 = new CountDownLatch(c1); + preStateChangeLatch2 = new CountDownLatch(c1); + preStateChangeCount1 = 0; + preStateChangeCount2 = 0; + postStateChangeLatch1 = new CountDownLatch(c1); + postStateChangeLatch2 = new CountDownLatch(c1); + postStateChangeCount1 = 0; + postStateChangeCount2 = 0; + preStateChangeStates1.clear(); + preStateChangeStates2.clear(); + postStateChangeStates1.clear(); + postStateChangeStates2.clear(); } @Override diff --git a/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/persist/PersistStateMachineHandler.java b/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/persist/PersistStateMachineHandler.java index 1f709781..bafa9c20 100644 --- a/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/persist/PersistStateMachineHandler.java +++ b/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/persist/PersistStateMachineHandler.java @@ -115,12 +115,6 @@ public class PersistStateMachineHandler extends LifecycleObjectSupport { private class PersistingStateChangeInterceptor extends StateMachineInterceptorAdapter { - @Override - public void preStateChange(State state, Message message, - Transition transition, StateMachine stateMachine) { - listeners.onPersist(state, message, transition, stateMachine); - } - @Override public void preStateChange(State state, Message message, Transition transition, StateMachine stateMachine, diff --git a/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/tasks/TasksHandler.java b/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/tasks/TasksHandler.java index 8a71f44f..91905497 100644 --- a/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/tasks/TasksHandler.java +++ b/spring-statemachine-recipes/src/main/java/org/springframework/statemachine/recipes/tasks/TasksHandler.java @@ -850,7 +850,8 @@ public class TasksHandler { @Override public void preStateChange(State state, Message message, - Transition transition, StateMachine stateMachine) { + Transition transition, StateMachine stateMachine, + StateMachine rootStateMachine) { // skip all other pseudostates than initial if (state == null || (state.getPseudoState() != null && state.getPseudoState().getKind() != PseudoStateKind.INITIAL)) { @@ -878,13 +879,6 @@ public class TasksHandler { throw new StateMachineException("Error persisting", e); } } - - @Override - public void preStateChange(State state, Message message, - Transition transition, StateMachine stateMachine, - StateMachine rootStateMachine) { - preStateChange(state, message, transition, stateMachine); - } } /**