From 5a779c885814245edfaf170569a97c503aecd67f Mon Sep 17 00:00:00 2001 From: Janne Valkealahti Date: Wed, 28 Sep 2016 15:40:26 +0100 Subject: [PATCH] Backport: Refactor join handling - Backporting #235 #237 - Refactor how join is handled - Add new StateListener for listening entry/exit per State. - Transition is now passed to next guy from a join state. - A lot of changes in core machine/factory/executor to support this new join handling. - Fixes #261 --- .../config/AbstractStateMachineFactory.java | 46 ++- .../statemachine/state/AbstractState.java | 13 + .../state/CompositeStateListener.java | 47 +++ .../statemachine/state/JoinPseudoState.java | 149 +++++-- .../statemachine/state/RegionState.java | 11 +- .../statemachine/state/State.java | 13 + .../statemachine/state/StateListener.java | 43 ++ .../statemachine/state/StateMachineState.java | 2 + .../support/AbstractStateMachine.java | 45 +- .../support/DefaultStateMachineExecutor.java | 31 +- .../support/StateMachineUtils.java | 8 + .../statemachine/state/JoinStateTests.java | 191 ++++++++- .../recipes/tasks/TasksHandler.java | 4 +- .../uml/UmlStateMachineModelFactoryTests.java | 51 +++ .../statemachine/uml/multijoin-forkjoin.di | 2 + .../uml/multijoin-forkjoin.notation | 383 ++++++++++++++++++ .../statemachine/uml/multijoin-forkjoin.uml | 63 +++ 17 files changed, 1021 insertions(+), 81 deletions(-) create mode 100644 spring-statemachine-core/src/main/java/org/springframework/statemachine/state/CompositeStateListener.java create mode 100644 spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateListener.java create mode 100644 spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.di create mode 100644 spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.notation create mode 100644 spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.uml diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/AbstractStateMachineFactory.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/AbstractStateMachineFactory.java index f55f19c3..78b0c53f 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/AbstractStateMachineFactory.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/AbstractStateMachineFactory.java @@ -61,6 +61,7 @@ import org.springframework.statemachine.state.ExitPseudoState; import org.springframework.statemachine.state.ForkPseudoState; import org.springframework.statemachine.state.HistoryPseudoState; import org.springframework.statemachine.state.JoinPseudoState; +import org.springframework.statemachine.state.JoinPseudoState.JoinStateData; import org.springframework.statemachine.state.JunctionPseudoState; import org.springframework.statemachine.state.JunctionPseudoState.JunctionStateData; import org.springframework.statemachine.state.PseudoState; @@ -631,34 +632,24 @@ public abstract class AbstractStateMachineFactory extends LifecycleObjectS joins.add(stateMap.get(fs)); } } - S ss = null; + + List> joinTargets = new ArrayList>(); Collection> transitions = stateMachineTransitions.getTransitions(); for (TransitionData tt : transitions) { if (tt.getSource() == s) { - ss = tt.getTarget(); - break; + StateHolder holder = new StateHolder(stateMap.get(tt.getTarget())); + if (holder.getState() == null) { + holderMap.put(tt.getTarget(), holder); + } + joinTargets.add(new JoinStateData(holder, tt.getGuard())); } } - StateHolder holder = new StateHolder(stateMap.get(ss)); - if (holder.getState() == null) { - holderMap.put(ss, holder); - } - JoinPseudoState pseudoState = new JoinPseudoState(joins, holder); + JoinPseudoState pseudoState = new JoinPseudoState(joins, joinTargets); + state = buildStateInternal(stateData.getState(), stateData.getDeferred(), stateData.getEntryActions(), stateData.getExitActions(), pseudoState); states.add(state); stateMap.put(stateData.getState(), state); - - // find joins sources and associate - for (Entry> e : stateMap.entrySet()) { - State value = e.getValue(); - if (value.isOrthogonal()) { - Collection> states2 = value.getStates(); - if (states2.containsAll(joins)) { - ((RegionState)value).setJoin(pseudoState); - } - } - } } } @@ -706,6 +697,23 @@ public abstract class AbstractStateMachineFactory extends LifecycleObjectS } } + if (stateMachineTransitions.getJoins() != null) { + for (Entry> entry : stateMachineTransitions.getJoins().entrySet()) { + if (stateMap.get(entry.getKey()) != null) { + List entryList = entry.getValue(); + for (S entryState : entryList) { + State source = stateMap.get(entryState); + if (source != null && !source.isOrthogonal()) { + State target = stateMap.get(entry.getKey()); + DefaultExternalTransition transition = new DefaultExternalTransition( + source, target, null, null, null, null, null); + transitions.add(transition); + } + } + } + } + } + Transition initialTransition = new InitialTransition(initialState, initialAction); StateMachine machine = buildStateMachineInternal(states, transitions, initialState, initialTransition, null, defaultExtendedState, historyState, contextEvents, beanFactory, taskExecutor, taskScheduler, diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractState.java index 3c662449..a38ad5c6 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/AbstractState.java @@ -44,6 +44,7 @@ public abstract class AbstractState implements State { private final Collection> regions = new ArrayList>(); private final StateMachine submachine; private List> triggers = new ArrayList>(); + private final CompositeStateListener stateListener = new CompositeStateListener(); /** * Instantiates a new abstract state. @@ -162,6 +163,7 @@ public abstract class AbstractState implements State { @Override public void exit(StateContext context) { + stateListener.onExit(context); for (Trigger trigger : triggers) { trigger.disarm(); } @@ -169,6 +171,7 @@ public abstract class AbstractState implements State { @Override public void entry(StateContext context) { + stateListener.onEntry(context); for (Trigger trigger : triggers) { trigger.arm(); } @@ -225,6 +228,16 @@ public abstract class AbstractState implements State { return submachine != null; } + @Override + public void addStateListener(StateListener listener) { + stateListener.register(listener); + } + + @Override + public void removeStateListener(StateListener listener) { + stateListener.unregister(listener); + } + /** * Gets the submachine. * diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/CompositeStateListener.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/CompositeStateListener.java new file mode 100644 index 00000000..4a739303 --- /dev/null +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/CompositeStateListener.java @@ -0,0 +1,47 @@ +/* + * Copyright 2016 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.statemachine.state; + +import java.util.Iterator; + +import org.springframework.statemachine.StateContext; +import org.springframework.statemachine.listener.AbstractCompositeListener; + +/** + * Composite state listener. + * + * @author Janne Valkealahti + * + * @param the type of state + * @param the type of event + */ +public class CompositeStateListener extends AbstractCompositeListener> + implements StateListener { + + @Override + public void onEntry(StateContext context) { + for (Iterator> iterator = getListeners().reverse(); iterator.hasNext();) { + iterator.next().onEntry(context); + } + } + + @Override + public void onExit(StateContext context) { + for (Iterator> iterator = getListeners().reverse(); iterator.hasNext();) { + iterator.next().onExit(context); + } + } +} 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 929f6daf..8e130523 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 @@ -18,9 +18,12 @@ package org.springframework.statemachine.state; import java.util.ArrayList; import java.util.List; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; import org.springframework.statemachine.StateContext; -import org.springframework.statemachine.listener.StateMachineListenerAdapter; +import org.springframework.statemachine.guard.Guard; import org.springframework.statemachine.state.PseudoStateContext.PseudoAction; +import org.springframework.statemachine.support.StateMachineUtils; import org.springframework.util.Assert; /** @@ -33,36 +36,42 @@ import org.springframework.util.Assert; */ public class JoinPseudoState extends AbstractPseudoState { + private final static Log log = LogFactory.getLog(JoinPseudoState.class); private final List> joins; - private volatile JoinTracker tracker; - private final StateHolder state; + private final JoinTracker tracker; + private final List> joinTargets; /** * Instantiates a new join pseudo state. * * @param joins the joins - * @param state the holder for target state + * @param joinTargets the target states */ - public JoinPseudoState(List> joins, StateHolder state) { + public JoinPseudoState(List> joins, List> joinTargets) { super(PseudoStateKind.JOIN); - Assert.notNull(state, "Holder must be set"); this.joins = joins; - this.state = state; + this.joinTargets = joinTargets; + this.tracker = new JoinTracker(); } @Override public State entry(StateContext context) { - tracker = new JoinTracker(this, new ArrayList>(joins)); - context.getStateMachine().addStateListener(tracker); - return state.getState(); + if (!tracker.isNotified()) { + return null; + } + State s = null; + for (JoinStateData c : joinTargets) { + s = c.getState(); + if (c.guard != null && evaluateInternal(c.guard, context)) { + break; + } + } + return s; } @Override public void exit(StateContext context) { - if (context != null) { - context.getStateMachine().removeStateListener(tracker); - } - tracker = null; + tracker.reset(); } /** @@ -74,28 +83,112 @@ public class JoinPseudoState extends AbstractPseudoState { return joins; } - private class JoinTracker extends StateMachineListenerAdapter { + private boolean evaluateInternal(Guard guard, StateContext context) { + try { + return guard.evaluate(context); + } catch (Throwable t) { + log.warn("Deny guard due to throw as GUARD should not error", t); + return false; + } + } + + private class JoinTracker { - private final PseudoState pseudoState; private final List> track; private volatile boolean notified = false; - public JoinTracker(PseudoState pseudoState, List> track) { - this.pseudoState = pseudoState; - this.track = track; - } + public JoinTracker() { + this.track = new ArrayList>(joins); + for (State tt : joins) { + final State t = tt; + t.addStateListener(new StateListener() { - @Override - public synchronized void stateChanged(State from, State to) { - if (!notified && track.size() > 0) { - track.remove(to); - if (track.size() == 0) { - notified = true; - notifyContext(new DefaultPseudoStateContext(pseudoState, PseudoAction.JOIN_COMPLETED)); - } + @Override + public void onEntry(StateContext context) { + if (StateMachineUtils.isPseudoState(context.getTransition().getTarget(), PseudoStateKind.END)) { + if (!notified && track.size() > 0) { + track.remove(t); + if (track.size() == 0) { + notified = true; + notifyContext(new DefaultPseudoStateContext(JoinPseudoState.this, PseudoAction.JOIN_COMPLETED)); + } + } + } + } + + @Override + public void onExit(StateContext context) { + if (!notified && track.size() > 0) { + track.remove(t); + if (track.size() == 0) { + notified = true; + notifyContext(new DefaultPseudoStateContext(JoinPseudoState.this, PseudoAction.JOIN_COMPLETED)); + } + } + } + }); } } + void reset() { + track.clear(); + track.addAll(joins); + notified = false; + } + + public boolean isNotified() { + return notified; + } } + /** + * Data class wrapping join {@link State} and {@link Guard} + * together. + * + * @param the type of state + * @param the type of event + */ + public static class JoinStateData { + private final StateHolder state; + private final Guard guard; + + /** + * Instantiates a new join state data. + * + * @param state the state holder + * @param guard the guard + */ + public JoinStateData(StateHolder state, Guard guard) { + Assert.notNull(state, "Holder must be set"); + this.state = state; + this.guard = guard; + } + + /** + * Gets the state holder. + * + * @return the state holder + */ + public StateHolder getStateHolder() { + return state; + } + + /** + * Gets the state. + * + * @return the state + */ + public State getState() { + return state.getState(); + } + + /** + * Gets the guard. + * + * @return the guard + */ + public Guard getGuard() { + return guard; + } + } } diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/RegionState.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/RegionState.java index cc54b650..1c07afec 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/RegionState.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/RegionState.java @@ -34,8 +34,6 @@ import org.springframework.statemachine.support.StateMachineUtils; */ public class RegionState extends AbstractState { - private JoinPseudoState join; - /** * Instantiates a new region state. * @@ -129,6 +127,7 @@ public class RegionState extends AbstractState { @Override public void exit(StateContext context) { + super.exit(context); for (Region region : getRegions()) { if (region.getState() != null) { region.getState().exit(context); @@ -145,9 +144,7 @@ public class RegionState extends AbstractState { @Override public void entry(StateContext context) { - if (join != null) { - join.entry(context); - } + super.entry(context); Collection> actions = getEntryActions(); if (actions != null) { for (Action action : actions) { @@ -201,10 +198,6 @@ public class RegionState extends AbstractState { return states; } - public void setJoin(JoinPseudoState join) { - this.join = join; - } - @Override public String toString() { return "RegionState [getIds()=" + getIds() + ", getClass()=" + getClass() + ", hashCode()=" + hashCode() diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/State.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/State.java index ec0babc1..339941f7 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/State.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/State.java @@ -147,4 +147,17 @@ public interface State { */ boolean isSubmachineState(); + /** + * Adds the state listener. + * + * @param listener the listener + */ + void addStateListener(StateListener listener); + + /** + * Removes the state listener. + * + * @param listener the listener + */ + void removeStateListener(StateListener listener); } diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateListener.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateListener.java new file mode 100644 index 00000000..48c6c166 --- /dev/null +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/state/StateListener.java @@ -0,0 +1,43 @@ +/* + * Copyright 2016 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.statemachine.state; + +import org.springframework.statemachine.StateContext; + +/** + * {@code StateListener} for various state events. + * + * @author Janne Valkealahti + * + * @param the type of state + * @param the type of event + */ +public interface StateListener { + + /** + * Called when {@link State} want to notify of its entry. + * + * @param context the state context + */ + void onEntry(StateContext context); + + /** + * Called when {@link State} want to notify of its exit. + * + * @param context the state context + */ + void onExit(StateContext context); +} 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 a60a7b53..957fb0b4 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 @@ -136,6 +136,7 @@ public class StateMachineState extends AbstractState { @Override public void exit(StateContext context) { + super.exit(context); // don't stop if it looks like we're coming back // stop would cause start with entry which would // enable default transition and state @@ -156,6 +157,7 @@ public class StateMachineState extends AbstractState { @Override public void entry(final StateContext context) { + super.entry(context); Collection> actions = getEntryActions(); if (actions != null && !isLocal(context)) { for (Action action : actions) { 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 93bd3990..3333a35c 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 @@ -45,7 +45,6 @@ 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; @@ -286,11 +285,15 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo // TODO: fix above stateContext as it's not used notifyTransitionStart(buildStateContext(Stage.TRANSITION_START, message, t, getRelayStateMachine())); notifyTransition(buildStateContext(Stage.TRANSITION, message, t, getRelayStateMachine())); - if (t.getKind() == TransitionKind.INITIAL) { - switchToState(t.getTarget(), message, t, getRelayStateMachine()); - notifyStateMachineStarted(buildStateContext(Stage.STATEMACHINE_START, message, t, getRelayStateMachine())); - } else if (t.getKind() != TransitionKind.INTERNAL) { - switchToState(t.getTarget(), message, t, getRelayStateMachine()); + if (t.getTarget().getPseudoState() != null && t.getTarget().getPseudoState().getKind() == PseudoStateKind.JOIN) { + exitFromState(t.getSource(), message, t, getRelayStateMachine()); + } else { + if (t.getKind() == TransitionKind.INITIAL) { + switchToState(t.getTarget(), message, t, getRelayStateMachine()); + notifyStateMachineStarted(buildStateContext(Stage.STATEMACHINE_START, message, t, getRelayStateMachine())); + } else if (t.getKind() != TransitionKind.INTERNAL) { + switchToState(t.getTarget(), message, t, getRelayStateMachine()); + } } // TODO: looks like events should be called here and anno processing earlier notifyTransitionEnd(buildStateContext(Stage.TRANSITION_END, message, t, getRelayStateMachine())); @@ -784,25 +787,29 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo @Override public void onContext(PseudoStateContext context) { PseudoState pseudoState = context.getPseudoState(); - if (pseudoState.getKind() == PseudoStateKind.JOIN) { - List> joins = ((JoinPseudoState)context.getPseudoState()).getJoins(); - for (State join : joins) { - exitFromState(join, null, null, getRelayStateMachine()); - } - } - State toState = findStateWithPseudoState(pseudoState); + State toStateOrig = findStateWithPseudoState(pseudoState); StateContext stateContext = buildStateContext(Stage.STATE_EXIT, null, null, getRelayStateMachine()); + State toState = followLinkedPseudoStates(toStateOrig, stateContext); + // TODO: try to find matching transition based on direct link. + // should make this built-in in pseudostates + Transition transition = findTransition(toStateOrig, toState); + switchToState(toState, null, transition, getRelayStateMachine()); pseudoState.exit(stateContext); - toState = followLinkedPseudoStates(toState, stateContext); - // should figure out what transition to use as we pass null for now - // which is then expected in exitCurrentState - switchToState(toState, null, null, getRelayStateMachine()); } }); } } } + private Transition findTransition(State from, State to) { + for (Transition transition : transitions) { + if (transition.getSource() == from && transition.getTarget() == to) { + return transition; + } + } + return null; + } + private State findStateWithPseudoState(PseudoState pseudoState) { for (State s : states) { if (s.getPseudoState() == pseudoState) { @@ -984,9 +991,7 @@ public abstract class AbstractStateMachine extends StateMachineObjectSuppo exitFromState(r.getState(), message, transition, stateMachine, sources, targets); } } - if (transition == null) { - exitFromState(currentState, message, transition, stateMachine, sources, targets); - } + exitFromState(currentState, message, transition, stateMachine, sources, targets); } else { exitFromState(currentState, message, transition, stateMachine, sources, targets); } diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/DefaultStateMachineExecutor.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/DefaultStateMachineExecutor.java index fad850c5..bbb2277b 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/DefaultStateMachineExecutor.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/DefaultStateMachineExecutor.java @@ -19,12 +19,14 @@ import java.util.ArrayList; import java.util.Collection; import java.util.Collections; import java.util.HashMap; +import java.util.HashSet; import java.util.LinkedList; import java.util.List; import java.util.ListIterator; import java.util.Map; import java.util.Map.Entry; import java.util.Queue; +import java.util.Set; import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; @@ -36,9 +38,11 @@ import org.springframework.core.task.TaskExecutor; import org.springframework.messaging.Message; import org.springframework.messaging.MessageHeaders; import org.springframework.statemachine.StateContext; +import org.springframework.statemachine.StateContext.Stage; import org.springframework.statemachine.StateMachine; import org.springframework.statemachine.StateMachineSystemConstants; -import org.springframework.statemachine.StateContext.Stage; +import org.springframework.statemachine.state.JoinPseudoState; +import org.springframework.statemachine.state.PseudoStateKind; import org.springframework.statemachine.state.State; import org.springframework.statemachine.transition.Transition; import org.springframework.statemachine.trigger.DefaultTriggerContext; @@ -172,6 +176,9 @@ public class DefaultStateMachineExecutor extends LifecycleObjectSupport im interceptors.add(interceptor); } + private final Set> joinSyncTransitions = new HashSet<>(); + private final Set> joinSyncStates = new HashSet<>(); + private boolean handleTriggerTrans(List> trans, Message queuedMessage) { boolean transit = false; for (Transition t : trans) { @@ -190,6 +197,28 @@ public class DefaultStateMachineExecutor extends LifecycleObjectSupport im continue; } + // special handling of join + if (StateMachineUtils.isPseudoState(t.getTarget(), PseudoStateKind.JOIN)) { + if (joinSyncStates.isEmpty()) { + List> joins = ((JoinPseudoState)t.getTarget().getPseudoState()).getJoins(); + joinSyncStates.addAll(joins); + } + joinSyncTransitions.add(t); + boolean removed = joinSyncStates.remove(t.getSource()); + boolean joincomplete = removed & joinSyncStates.isEmpty(); + if (joincomplete) { + for (Transition tt : joinSyncTransitions) { + StateContext stateContext = buildStateContext(queuedMessage, tt, relayStateMachine); + tt.transit(stateContext); + stateMachineExecutorTransit.transit(tt, stateContext, queuedMessage); + } + joinSyncTransitions.clear(); + break; + } else { + continue; + } + } + StateContext stateContext = buildStateContext(queuedMessage, t, relayStateMachine); try { stateContext = interceptors.preTransition(stateContext); diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineUtils.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineUtils.java index 4ccf680b..3e669d0e 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineUtils.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/support/StateMachineUtils.java @@ -95,6 +95,14 @@ public abstract class StateMachineUtils { } } + public static boolean isPseudoState(State state, PseudoStateKind kind) { + if (state != null) { + PseudoState pseudoState = state.getPseudoState(); + return pseudoState != null && pseudoState.getKind() == kind; + } + return false; + } + public static Collection toStringCollection(Collection collection) { Collection c = new ArrayList(); for (S item : collection) { 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 5e787513..d63de650 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,6 +16,7 @@ package org.springframework.statemachine.state; 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.junit.Assert.assertThat; @@ -24,18 +25,22 @@ import java.util.ArrayList; import java.util.List; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import org.junit.Test; import org.springframework.context.annotation.AnnotationConfigApplicationContext; import org.springframework.context.annotation.Configuration; +import org.springframework.messaging.Message; import org.springframework.statemachine.AbstractStateMachineTests; import org.springframework.statemachine.ObjectStateMachine; +import org.springframework.statemachine.StateMachine; import org.springframework.statemachine.StateMachineSystemConstants; 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.support.StateMachineInterceptorAdapter; import org.springframework.statemachine.transition.Transition; public class JoinStateTests extends AbstractStateMachineTests { @@ -105,16 +110,17 @@ public class JoinStateTests extends AbstractStateMachineTests { 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)); - listener.reset(3); + listener.reset(2); machine.sendEvent(TestEvents.E3); assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); - assertThat(listener.stateChangedCount, is(3)); + assertThat(listener.stateChangedCount, is(2)); assertThat(machine.getState().getIds(), contains(TestStates.S4)); } @@ -185,14 +191,128 @@ public class JoinStateTests extends AbstractStateMachineTests { assertThat(listener.stateChangedLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(listener.stateChangedCount, is(1)); - listener.reset(3); + 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 testMultiJoin1() throws Exception { + context.register(BaseConfig.class, Config3.class); + context.refresh(); + ObjectStateMachine machine = + context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, ObjectStateMachine.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(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 testMultiJoin2() throws Exception { + context.register(BaseConfig.class, Config3.class); + context.refresh(); + ObjectStateMachine machine = + context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, ObjectStateMachine.class); + TestListener listener = new TestListener(); + machine.addStateListener(listener); + listener.reset(1); + assertThat(machine, notNullValue()); + machine.start(); + machine.getExtendedState().getVariables().put("foo", "bar"); + 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(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.SF)); + } + + @Test + @SuppressWarnings("unchecked") + public void testInterceptorPostStateChangeTransitionNotNull() throws Exception { + context.register(BaseConfig.class, Config1.class); + context.refresh(); + ObjectStateMachine machine = + context.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, ObjectStateMachine.class); + TestListener listener = new TestListener(); + machine.addStateListener(listener); + + final AtomicBoolean nullCheck = new AtomicBoolean(false); + machine.addStateMachineInterceptor(new StateMachineInterceptorAdapter() { + @Override + public void postStateChange(State state, Message message, + Transition transition, StateMachine stateMachine) { + if (state.getId() == TestStates.S4) { + nullCheck.set(transition == null); + } + super.postStateChange(state, message, transition, stateMachine); + } + }); + + 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(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)); + assertThat("Interceptor postStateChange has null transition", nullCheck.get(), is(false)); + } + @Configuration @EnableStateMachine static class Config1 extends EnumStateMachineConfigurerAdapter { @@ -312,6 +432,71 @@ public class JoinStateTests extends AbstractStateMachineTests { } + @Configuration + @EnableStateMachine + 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.SF) + .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.SF) + .guardExpression("!extendedState.variables.isEmpty()") + .and() + .withExternal() + .source(TestStates.S3) + .target(TestStates.S4) + .guardExpression("extendedState.variables.isEmpty()") + .and() + .withExternal() + .source(TestStates.S4) + .target(TestStates.SI) + .event(TestEvents.E4); + } + + } + private static class TestListener extends StateMachineListenerAdapter { volatile CountDownLatch stateChangedLatch = new CountDownLatch(1); 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 b8d0c989..bcf60f7c 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 @@ -272,7 +272,9 @@ public class TasksHandler { .initial(initial) .state(task, runnableAction(node.getData().runnable, node.getData().id.toString()), null); - joinStates.add(task); + if (node.getChildren().isEmpty()) { + joinStates.add(task); + } stateMachineTransitionConfigurer .withExternal() diff --git a/spring-statemachine-uml/src/test/java/org/springframework/statemachine/uml/UmlStateMachineModelFactoryTests.java b/spring-statemachine-uml/src/test/java/org/springframework/statemachine/uml/UmlStateMachineModelFactoryTests.java index 12d48132..12464673 100644 --- a/spring-statemachine-uml/src/test/java/org/springframework/statemachine/uml/UmlStateMachineModelFactoryTests.java +++ b/spring-statemachine-uml/src/test/java/org/springframework/statemachine/uml/UmlStateMachineModelFactoryTests.java @@ -339,6 +339,39 @@ public class UmlStateMachineModelFactoryTests extends AbstractUmlTests { assertThat(stateMachine.getState().getIds(), containsInAnyOrder("SF")); } + @Test + @SuppressWarnings("unchecked") + public void testMultiJoinForkJoin1() { + context.register(Config20.class); + context.refresh(); + StateMachine stateMachine = context.getBean(StateMachine.class); + stateMachine.start(); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("SI")); + stateMachine.sendEvent("E1"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("S2", "S20", "S30")); + stateMachine.sendEvent("E2"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("S2", "S21", "S30")); + stateMachine.sendEvent("E3"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("S4")); + } + + @Test + @SuppressWarnings("unchecked") + public void testMultiJoinForkJoin2() { + context.register(Config20.class); + context.refresh(); + StateMachine stateMachine = context.getBean(StateMachine.class); + stateMachine.start(); + stateMachine.getExtendedState().getVariables().put("foo", "bar"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("SI")); + stateMachine.sendEvent("E1"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("S2", "S20", "S30")); + stateMachine.sendEvent("E2"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("S2", "S21", "S30")); + stateMachine.sendEvent("E3"); + assertThat(stateMachine.getState().getIds(), containsInAnyOrder("SF")); + } + @Test @SuppressWarnings("unchecked") public void testSimpleHistoryShallow() { @@ -984,6 +1017,24 @@ public class UmlStateMachineModelFactoryTests extends AbstractUmlTests { } } + @Configuration + @EnableStateMachine + public static class Config20 extends StateMachineConfigurerAdapter { + + @Override + public void configure(StateMachineModelConfigurer model) throws Exception { + model + .withModel() + .factory(modelFactory()); + } + + @Bean + public StateMachineModelFactory modelFactory() { + Resource model = new ClassPathResource("org/springframework/statemachine/uml/multijoin-forkjoin.uml"); + return new UmlStateMachineModelFactory(model); + } + } + public static class LatchAction implements Action { CountDownLatch latch = new CountDownLatch(1); @Override diff --git a/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.di b/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.di new file mode 100644 index 00000000..bf9abab3 --- /dev/null +++ b/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.di @@ -0,0 +1,2 @@ + + diff --git a/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.notation b/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.notation new file mode 100644 index 00000000..cc223e30 --- /dev/null +++ b/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.notation @@ -0,0 +1,383 @@ + + + + + + + + + +
+ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
+ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
+ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.uml b/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.uml new file mode 100644 index 00000000..7f9ec2e4 --- /dev/null +++ b/spring-statemachine-uml/src/test/resources/org/springframework/statemachine/uml/multijoin-forkjoin.uml @@ -0,0 +1,63 @@ + + + + + + + + + + + + + + spel + !extendedState.variables.isEmpty() + + + + + + + spel + extendedState.variables.isEmpty() + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +