From 06f839fb71f70c973d900e0082842c240afe1355 Mon Sep 17 00:00:00 2001 From: Janne Valkealahti Date: Fri, 22 Apr 2016 15:04:23 +0100 Subject: [PATCH] Guard should not break with Throwable - Now catching all exceptions and throwables with evaluate and return false in those cases. - Fixes #206 --- .../statemachine/state/ChoicePseudoState.java | 14 ++- .../state/JunctionPseudoState.java | 14 ++- .../transition/AbstractTransition.java | 11 +- .../statemachine/guard/GuardTests.java | 102 +++++++++++++++--- 4 files changed, 124 insertions(+), 17 deletions(-) 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 4c8645c9..28a42d6d 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 @@ -17,6 +17,8 @@ package org.springframework.statemachine.state; 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.guard.Guard; import org.springframework.util.Assert; @@ -31,6 +33,7 @@ import org.springframework.util.Assert; */ public class ChoicePseudoState implements PseudoState { + private final static Log log = LogFactory.getLog(ChoicePseudoState.class); private final List> choices; /** @@ -52,7 +55,7 @@ public class ChoicePseudoState implements PseudoState { State s = null; for (ChoiceStateData c : choices) { s = c.getState(); - if (c.guard != null && c.guard.evaluate(context)) { + if (c.guard != null && evaluateInternal(c.guard, context)) { break; } } @@ -67,6 +70,15 @@ public class ChoicePseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { } + 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; + } + } + /** * Data class wrapping choice {@link State} and {@link Guard} * together. 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 dada89e1..ca092f9d 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 @@ -17,6 +17,8 @@ package org.springframework.statemachine.state; 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.guard.Guard; import org.springframework.util.Assert; @@ -31,6 +33,7 @@ import org.springframework.util.Assert; */ public class JunctionPseudoState implements PseudoState { + private final static Log log = LogFactory.getLog(JunctionPseudoState.class); private final List> junctions; /** @@ -52,7 +55,7 @@ public class JunctionPseudoState implements PseudoState { State s = null; for (JunctionStateData j : junctions) { s = j.getState(); - if (j.guard != null && j.guard.evaluate(context)) { + if (j.guard != null && evaluateInternal(j.guard, context)) { break; } } @@ -67,6 +70,15 @@ public class JunctionPseudoState implements PseudoState { public void addPseudoStateListener(PseudoStateListener listener) { } + 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; + } + } + /** * Data class wrapping choice {@link State} and {@link Guard} * together. diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/transition/AbstractTransition.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/transition/AbstractTransition.java index 900e3e8d..517e90c9 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/transition/AbstractTransition.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/transition/AbstractTransition.java @@ -17,6 +17,8 @@ package org.springframework.statemachine.transition; import java.util.Collection; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; import org.springframework.statemachine.StateContext; import org.springframework.statemachine.action.Action; import org.springframework.statemachine.guard.Guard; @@ -35,6 +37,8 @@ import org.springframework.util.Assert; */ public abstract class AbstractTransition implements Transition { + private final static Log log = LogFactory.getLog(AbstractTransition.class); + private final State source; private final State target; @@ -90,7 +94,12 @@ public abstract class AbstractTransition implements Transition { @Override public boolean transit(StateContext context) { if (guard != null) { - if (!guard.evaluate(context)) { + try { + if (!guard.evaluate(context)) { + return false; + } + } catch (Throwable t) { + log.warn("Deny guard due to throw as GUARD should not error", t); return false; } } diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/guard/GuardTests.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/guard/GuardTests.java index 9d0ace3b..0b164be0 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/guard/GuardTests.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/guard/GuardTests.java @@ -33,6 +33,7 @@ import org.springframework.statemachine.AbstractStateMachineTests.TestEvents; import org.springframework.statemachine.AbstractStateMachineTests.TestGuard; import org.springframework.statemachine.AbstractStateMachineTests.TestStates; import org.springframework.statemachine.ObjectStateMachine; +import org.springframework.statemachine.StateContext; import org.springframework.statemachine.StateMachineSystemConstants; import org.springframework.statemachine.config.EnableStateMachine; import org.springframework.statemachine.config.EnumStateMachineConfigurerAdapter; @@ -41,7 +42,7 @@ import org.springframework.statemachine.config.builders.StateMachineTransitionCo /** * Tests for state machine guards. - * + * * @author Janne Valkealahti * */ @@ -57,12 +58,12 @@ public class GuardTests { TestAction testAction = ctx.getBean("testAction", TestAction.class); assertThat(testGuard, notNullValue()); assertThat(testAction, notNullValue()); - + machine.start(); - machine.sendEvent(TestEvents.E1); + machine.sendEvent(TestEvents.E1); assertThat(testGuard.onEvaluateLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(testAction.onExecuteLatch.await(2, TimeUnit.SECONDS), is(true)); - + ctx.close(); } @@ -76,18 +77,35 @@ public class GuardTests { TestAction testAction = ctx.getBean("testAction", TestAction.class); assertThat(testGuard, notNullValue()); assertThat(testAction, notNullValue()); - + machine.start(); assertThat(machine.getState().getIds(), contains(TestStates.S1)); - - machine.sendEvent(TestEvents.E1); + + machine.sendEvent(TestEvents.E1); assertThat(testGuard.onEvaluateLatch.await(2, TimeUnit.SECONDS), is(true)); assertThat(testAction.onExecuteLatch.await(2, TimeUnit.SECONDS), is(false)); assertThat(machine.getState().getIds(), contains(TestStates.S1)); - + ctx.close(); } - + + @SuppressWarnings({ "unchecked" }) + @Test + public void testGuardThrows() throws Exception { + AnnotationConfigApplicationContext ctx = new AnnotationConfigApplicationContext(Config3.class); + ObjectStateMachine machine = + ctx.getBean(StateMachineSystemConstants.DEFAULT_ID_STATEMACHINE, ObjectStateMachine.class); + machine.start(); + assertThat(machine.getState().getIds(), contains(TestStates.S1)); + machine.sendEvent(TestEvents.E1); + assertThat(machine.getState().getIds(), contains(TestStates.S1)); + machine.sendEvent(TestEvents.E2); + assertThat(machine.getState().getIds(), contains(TestStates.S1)); + machine.sendEvent(TestEvents.E3); + assertThat(machine.getState().getIds(), contains(TestStates.S2)); + ctx.close(); + } + @Configuration @EnableStateMachine public static class Config1 extends EnumStateMachineConfigurerAdapter { @@ -116,12 +134,12 @@ public class GuardTests { public TestAction testAction() { return new TestAction(); } - + @Bean public TestGuard testGuard() { return new TestGuard(true); } - + @Bean public TaskExecutor taskExecutor() { return new SyncTaskExecutor(); @@ -152,12 +170,12 @@ public class GuardTests { .action(testAction()) .guard(testGuard()); } - + @Bean public TestGuard testGuard() { return new TestGuard(false); } - + @Bean public TestAction testAction() { return new TestAction(); @@ -169,5 +187,61 @@ public class GuardTests { } } - + + @Configuration + @EnableStateMachine + public static class Config3 extends EnumStateMachineConfigurerAdapter { + + @Override + public void configure(StateMachineStateConfigurer states) throws Exception { + states + .withStates() + .initial(TestStates.S1) + .state(TestStates.S1) + .state(TestStates.S2); + } + + @Override + public void configure(StateMachineTransitionConfigurer transitions) throws Exception { + transitions + .withExternal() + .source(TestStates.S1) + .target(TestStates.S2) + .event(TestEvents.E1) + .guard(testGuard1()) + .and() + .withExternal() + .source(TestStates.S1) + .target(TestStates.S2) + .event(TestEvents.E2) + .guard(testGuard2()) + .and() + .withExternal() + .source(TestStates.S1) + .target(TestStates.S2) + .event(TestEvents.E3); + } + + @Bean + public Guard testGuard1() { + return new Guard() { + + @Override + public boolean evaluate(StateContext context) { + throw new RuntimeException(); + } + }; + } + + @Bean + public Guard testGuard2() { + return new Guard() { + + @Override + public boolean evaluate(StateContext context) { + throw new Error(); + } + }; + } + } }