diff --git a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/configuration/StateMachineFactoryConfiguration.java b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/configuration/StateMachineFactoryConfiguration.java index 0b3b6a42..be3e6d7d 100644 --- a/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/configuration/StateMachineFactoryConfiguration.java +++ b/spring-statemachine-core/src/main/java/org/springframework/statemachine/config/configuration/StateMachineFactoryConfiguration.java @@ -40,6 +40,7 @@ import org.springframework.statemachine.config.builders.StateMachineStates; import org.springframework.statemachine.config.builders.StateMachineTransitions; import org.springframework.statemachine.config.common.annotation.AbstractImportingAnnotationConfiguration; import org.springframework.statemachine.config.common.annotation.AnnotationConfigurer; +import org.springframework.util.ClassUtils; @Configuration public class StateMachineFactoryConfiguration, E extends Enum> extends @@ -56,6 +57,7 @@ public class StateMachineFactoryConfiguration, E extends Enum< EnableStateMachineFactory.class.getName(), false)); Boolean contextEvents = attributes.getBoolean("contextEvents"); beanDefinitionBuilder.addConstructorArgValue(builder); + beanDefinitionBuilder.addConstructorArgValue(importingClassMetadata.getClassName()); beanDefinitionBuilder.addConstructorArgValue(contextEvents); return beanDefinitionBuilder.getBeanDefinition(); } @@ -72,18 +74,16 @@ public class StateMachineFactoryConfiguration, E extends Enum< FactoryBean>, BeanFactoryAware, InitializingBean { private final StateMachineConfigBuilder builder; - private List, StateMachineConfigBuilder>> configurers; - private BeanFactory beanFactory; - private StateMachineFactory stateMachineFactory; - + private String clazzName; private Boolean contextEvents; @SuppressWarnings("unused") - public StateMachineFactoryDelegatingFactoryBean(StateMachineConfigBuilder builder, Boolean contextEvents) { + public StateMachineFactoryDelegatingFactoryBean(StateMachineConfigBuilder builder, String clazzName, Boolean contextEvents) { this.builder = builder; + this.clazzName = clazzName; this.contextEvents = contextEvents; } @@ -105,7 +105,10 @@ public class StateMachineFactoryConfiguration, E extends Enum< @Override public void afterPropertiesSet() throws Exception { for (AnnotationConfigurer, StateMachineConfigBuilder> configurer : configurers) { - builder.apply(configurer); + Class clazz = configurer.getClass(); + if (ClassUtils.getUserClass(clazz).getName().equals(clazzName)) { + builder.apply(configurer); + } } StateMachineConfig stateMachineConfig = builder.getOrBuild(); StateMachineTransitions stateMachineTransitions = stateMachineConfig.getTransitions(); diff --git a/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineFactoryTests.java b/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineFactoryTests.java index 19cc3bcc..b3af6f02 100644 --- a/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineFactoryTests.java +++ b/spring-statemachine-core/src/test/java/org/springframework/statemachine/StateMachineFactoryTests.java @@ -86,6 +86,38 @@ public class StateMachineFactoryTests extends AbstractStateMachineTests { assertThat(((SmartLifecycle)machine).isRunning(), is(false)); } + @SuppressWarnings({ "unchecked" }) + @Test + public void testCustomNamedFactory() { + context.register(Config4.class); + context.refresh(); + StateMachineFactory stateMachineFactory = + context.getBean("factory1", ObjectStateMachineFactory.class); + StateMachine machine = stateMachineFactory.getStateMachine(); + machine.start(); + + assertThat(machine.getState().getIds(), contains(TestStates.S1)); + } + + @SuppressWarnings({ "unchecked" }) + @Test + public void testMultipleCustomNamedFactories() { + context.register(Config4.class, Config5.class); + context.refresh(); + StateMachineFactory stateMachineFactory1 = + context.getBean("factory1", ObjectStateMachineFactory.class); + StateMachineFactory stateMachineFactory2 = + context.getBean("factory2", ObjectStateMachineFactory.class); + StateMachine machine1 = stateMachineFactory1.getStateMachine(); + StateMachine machine2 = stateMachineFactory2.getStateMachine(); + + machine1.start(); + machine2.start(); + + assertThat(machine1.getState().getIds(), contains(TestStates.S1)); + assertThat(machine2.getState().getIds(), contains(TestStates.S1)); + } + @Configuration @EnableStateMachineFactory static class Config1 extends EnumStateMachineConfigurerAdapter { @@ -157,4 +189,64 @@ public class StateMachineFactoryTests extends AbstractStateMachineTests { } + @Configuration + @EnableStateMachineFactory(name = "factory1") + public static class Config4 extends EnumStateMachineConfigurerAdapter { + + @Override + public void configure(StateMachineConfigurationConfigurer config) throws Exception { + config + .withConfiguration() + .autoStartup(false); + } + + @Override + public void configure(StateMachineStateConfigurer states) throws Exception { + states + .withStates() + .initial(TestStates.S1) + .states(EnumSet.allOf(TestStates.class)); + } + + @Override + public void configure(StateMachineTransitionConfigurer transitions) throws Exception { + transitions + .withExternal() + .source(TestStates.S1) + .target(TestStates.S2) + .event(TestEvents.E1); + } + + } + + @Configuration + @EnableStateMachineFactory(name = "factory2") + public static class Config5 extends EnumStateMachineConfigurerAdapter { + + @Override + public void configure(StateMachineConfigurationConfigurer config) throws Exception { + config + .withConfiguration() + .autoStartup(false); + } + + @Override + public void configure(StateMachineStateConfigurer states) throws Exception { + states + .withStates() + .initial(TestStates.S1) + .states(EnumSet.allOf(TestStates.class)); + } + + @Override + public void configure(StateMachineTransitionConfigurer transitions) throws Exception { + transitions + .withExternal() + .source(TestStates.S1) + .target(TestStates.S2) + .event(TestEvents.E1); + } + + } + }