diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java index 30ec2678..3f7be089 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/EnumStateMachineFactory.java @@ -37,6 +37,7 @@ import org.springframework.statemachine.state.PseudoStateKind; import org.springframework.statemachine.state.RegionState; import org.springframework.statemachine.state.State; import org.springframework.statemachine.state.StateMachineState; +import org.springframework.statemachine.support.DefaultExtendedState; import org.springframework.statemachine.support.LifecycleObjectSupport; import org.springframework.statemachine.support.tree.Tree; import org.springframework.statemachine.support.tree.Tree.Node; @@ -82,6 +83,10 @@ public class EnumStateMachineFactory, E extends Enum> exten @Override public StateMachine getStateMachine() { + + // shared + DefaultExtendedState defaultExtendedState = new DefaultExtendedState(); + StateMachine machine = null; // we store mappings from state id's to states which gets @@ -142,7 +147,7 @@ public class EnumStateMachineFactory, E extends Enum> exten Collection> stateDatas = popSameParents(stateStack); Collection> transitionsData = getTransitionData(iterator.hasNext(), stateDatas); - machine = buildMachine(machineMap, stateMap, stateDatas, transitionsData, getBeanFactory(), contextEvents); + machine = buildMachine(machineMap, stateMap, stateDatas, transitionsData, getBeanFactory(), contextEvents, defaultExtendedState); // TODO: last part in if feels a bit hack // (!peek.isInitial() && !machineMap.containsKey(peek.getParent())) if (peek.isInitial() || (!peek.isInitial() && !machineMap.containsKey(peek.getParent()))) { @@ -170,7 +175,7 @@ public class EnumStateMachineFactory, E extends Enum> exten } if (initials == 1) { - machine = buildMachine(machineMap, stateMap, stateDatas, transitionsData, getBeanFactory(), contextEvents); + machine = buildMachine(machineMap, stateMap, stateDatas, transitionsData, getBeanFactory(), contextEvents, defaultExtendedState); } } @@ -185,7 +190,8 @@ public class EnumStateMachineFactory, E extends Enum> exten RegionState rstate = new RegionState(null, regions); Collection> states = new ArrayList>(); states.add(rstate); - EnumStateMachine m = new EnumStateMachine(states, null, rstate, null); + EnumStateMachine m = new EnumStateMachine(states, null, rstate, + null, null, null, defaultExtendedState); if (contextEvents != null) { m.setContextEventsEnabled(contextEvents); } @@ -266,8 +272,9 @@ public class EnumStateMachineFactory, E extends Enum> exten private static , E extends Enum> StateMachine buildMachine( - Map> machineMap, Map> stateMap, Collection> stateDatas, - Collection> transitionsData, BeanFactory beanFactory, Boolean contextEvents) { + Map> machineMap, Map> stateMap, + Collection> stateDatas, Collection> transitionsData, + BeanFactory beanFactory, Boolean contextEvents, DefaultExtendedState defaultExtendedState) { State state = null; State initialState = null; Action initialAction = null; @@ -348,7 +355,7 @@ public class EnumStateMachineFactory, E extends Enum> exten } EnumStateMachine machine = new EnumStateMachine(states, transitions, initialState, - initialTransition, endState, null, null); + initialTransition, endState, null, defaultExtendedState); if (contextEvents != null) { machine.setContextEventsEnabled(contextEvents); } 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 b31fd89d..3f8e0511 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 @@ -223,7 +223,7 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport protected void doStart() { super.doStart(); registerTriggerListener(); - switchToState(initialState, initialEvent, null); + switchToState(initialState, initialEvent, null, this); // TODO: for now execute outside of switchToState if (initialTransition != null) { StateContext stateContext = new DefaultStateContext( @@ -322,15 +322,15 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport return false; } - private void switchToState(State state, Message event, Transition transition) { - setCurrentState(state, event, transition, true); + private void switchToState(State state, Message event, Transition transition, StateMachine stateMachine) { + setCurrentState(state, event, transition, true, stateMachine); // TODO: should handle triggerles transition some how differently for (Transition t : transitions) { State source = t.getSource(); State target = t.getTarget(); if (t.getTrigger() == null && source.equals(currentState)) { - switchToState(target, event, t); + switchToState(target, event, t, stateMachine); } } @@ -345,7 +345,7 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport return null; } - void setCurrentState(State state, Message event, Transition transition, boolean exit) { + void setCurrentState(State state, Message event, Transition transition, boolean exit, StateMachine stateMachine) { State findDeep = findDeepParent(state); boolean isTargetSubOf = false; if (transition != null) { @@ -357,49 +357,49 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport if (states.contains(state)) { if (exit) { - exitCurrentState(state, event, transition); + exitCurrentState(state, event, transition, stateMachine); } State notifyFrom = currentState; currentState = state; - entryToState(state, event, transition); + entryToState(state, event, transition, stateMachine); notifyStateChanged(notifyFrom, state); } else if (currentState != null && currentState.isSubmachineState()) { if (findDeep != null) { if (exit) { - exitCurrentState(state, event, transition); + exitCurrentState(state, event, transition, stateMachine); } if (currentState == findDeep) { StateMachine submachine = ((AbstractState)currentState).getSubmachine(); if (submachine.getState() == state) { if (currentState == findDeep) { if (isTargetSubOf) { - entryToState(currentState, event, transition); + entryToState(currentState, event, transition, stateMachine); } currentState = findDeep; - ((AbstractStateMachine)submachine).setCurrentState(state, event, transition, false); + ((AbstractStateMachine)submachine).setCurrentState(state, event, transition, false, stateMachine); return; } } } currentState = findDeep; - entryToState(currentState, event, transition); + entryToState(currentState, event, transition, stateMachine); StateMachine submachine = ((AbstractState)currentState).getSubmachine(); - ((AbstractStateMachine)submachine).setCurrentState(state, event, transition, false); + ((AbstractStateMachine)submachine).setCurrentState(state, event, transition, false, stateMachine); } } } - void exitCurrentState(State state, Message event, Transition transition) { + void exitCurrentState(State state, Message event, Transition transition, StateMachine stateMachine) { if (currentState == null) { return; } if (currentState.isSubmachineState()) { StateMachine submachine = ((AbstractState)currentState).getSubmachine(); - ((AbstractStateMachine)submachine).exitCurrentState(state, event, transition); - exitFromState(currentState, event, transition); + ((AbstractStateMachine)submachine).exitCurrentState(state, event, transition, stateMachine); + exitFromState(currentState, event, transition, stateMachine); } else { - exitFromState(currentState, event, transition); + exitFromState(currentState, event, transition, stateMachine); } } @@ -409,12 +409,12 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport return c.contains(right); } - private void exitFromState(State state, Message event, Transition transition) { + private void exitFromState(State state, Message event, Transition transition, StateMachine stateMachine) { if (state != null) { log.trace("Exit state=[" + state + "]"); MessageHeaders messageHeaders = event != null ? event.getHeaders() : new MessageHeaders( new HashMap()); - StateContext stateContext = new DefaultStateContext(messageHeaders, extendedState, transition, this); + StateContext stateContext = new DefaultStateContext(messageHeaders, extendedState, transition, stateMachine); State findDeep = findDeepParent(transition.getTarget()); boolean isTargetSubOfOtherState = findDeep != null && findDeep != currentState; @@ -439,12 +439,12 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport } } - private void entryToState(State state, Message event, Transition transition) { + private void entryToState(State state, Message event, Transition transition, StateMachine stateMachine) { if (state != null) { log.trace("Enter state=[" + state + "]"); MessageHeaders messageHeaders = event != null ? event.getHeaders() : new MessageHeaders( new HashMap()); - StateContext stateContext = new DefaultStateContext(messageHeaders, extendedState, transition, this); + StateContext stateContext = new DefaultStateContext(messageHeaders, extendedState, transition, stateMachine); if (transition != null) { State findDeep1 = findDeepParent(transition.getTarget()); @@ -591,7 +591,7 @@ public abstract class AbstractStateMachine extends LifecycleObjectSupport notifyTransitionStart(t); callHandlers(t.getSource(), t.getTarget(), queuedEvent); if (t.getKind() != TransitionKind.INTERNAL) { - switchToState(t.getTarget(), queuedEvent, t); + switchToState(t.getTarget(), queuedEvent, t, this); } notifyTransition(t); notifyTransitionEnd(t); diff --git a/spring-statemachine-samples/cdplayer/src/main/java/demo/cdplayer/Application.java b/spring-statemachine-samples/cdplayer/src/main/java/demo/cdplayer/Application.java index 15e8796b..8b956a5c 100644 --- a/spring-statemachine-samples/cdplayer/src/main/java/demo/cdplayer/Application.java +++ b/spring-statemachine-samples/cdplayer/src/main/java/demo/cdplayer/Application.java @@ -196,8 +196,8 @@ public class Application { @Override public void execute(StateContext context) { if (context.getTransition() != null - && context.getTransition().getSource().getId() == States.CLOSED - && context.getMessageHeader(Variables.CD) != null) { + && context.getTransition().getTarget().getId() == States.CLOSED + && context.getExtendedState().getVariables().get(Variables.CD) != null) { context.getStateMachine().sendEvent(Events.PLAY); } } diff --git a/spring-statemachine-samples/cdplayer/src/test/java/demo/cdplayer/CdPlayerTests.java b/spring-statemachine-samples/cdplayer/src/test/java/demo/cdplayer/CdPlayerTests.java index f79b1f38..86878dde 100644 --- a/spring-statemachine-samples/cdplayer/src/test/java/demo/cdplayer/CdPlayerTests.java +++ b/spring-statemachine-samples/cdplayer/src/test/java/demo/cdplayer/CdPlayerTests.java @@ -77,6 +77,18 @@ public class CdPlayerTests { assertLcdStatusContains("cd1"); } + @Test + public void testPlayWithCdLoadedDeckOpen() throws Exception { + listener.reset(3, 0, 0); + player.eject(); + player.load(library.getCollection().get(0)); + player.play(); + listener.stateChangedLatch.await(5, TimeUnit.SECONDS); + assertThat(listener.stateChangedCount, is(4)); + assertThat(machine.getState().getIds(), contains(States.BUSY, States.PLAYING)); + assertLcdStatusContains("cd1"); + } + @Test public void testPlayWithNoCdLoaded() throws Exception { listener.reset(0, 0, 0);