diff --git a/spring-integration-core/src/main/java/org/springframework/integration/router/AbstractMessageRouter.java b/spring-integration-core/src/main/java/org/springframework/integration/router/AbstractMessageRouter.java index 1eae37810a..5183f017ad 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/router/AbstractMessageRouter.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/router/AbstractMessageRouter.java @@ -21,6 +21,7 @@ import java.util.Arrays; import java.util.Collection; import java.util.List; import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; import org.springframework.beans.factory.BeanFactory; import org.springframework.core.convert.ConversionService; @@ -215,11 +216,14 @@ public abstract class AbstractMessageRouter extends AbstractMessageHandler { this.channelIdentifierMap = channelIdentifierMap; } - public void setChannelMapping(String channelIdentifier, String channelName){ + public synchronized void setChannelMapping(String channelIdentifier, String channelName){ + if (channelIdentifierMap == null){ + channelIdentifierMap = new ConcurrentHashMap(); + } this.channelIdentifierMap.put(channelIdentifier, channelName); } - public void removeChannelMapping(String channelIdentifier){ + public synchronized void removeChannelMapping(String channelIdentifier){ this.channelIdentifierMap.remove(channelIdentifier); } /** @@ -320,10 +324,4 @@ public abstract class AbstractMessageRouter extends AbstractMessageHandler { } } } - -// /** -// * Subclasses must implement this method to return the target channels for a given Message. -// */ -// protected abstract Collection determineTargetChannels(Message message); - } diff --git a/spring-integration-core/src/main/java/org/springframework/integration/router/PayloadTypeRouter.java b/spring-integration-core/src/main/java/org/springframework/integration/router/PayloadTypeRouter.java index 1c7e965fcb..3293db5845 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/router/PayloadTypeRouter.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/router/PayloadTypeRouter.java @@ -19,8 +19,11 @@ package org.springframework.integration.router; import java.util.Collections; import java.util.List; +import org.springframework.beans.factory.BeanFactory; import org.springframework.integration.Message; import org.springframework.integration.MessageChannel; +import org.springframework.integration.support.channel.BeanFactoryChannelResolver; +import org.springframework.util.Assert; import org.springframework.util.StringUtils; /** @@ -37,7 +40,7 @@ public class PayloadTypeRouter extends AbstractMessageRouter { Class firstInterfaceMatch = null; Class type = message.getPayload().getClass(); - while (type != null) { + while (type != null && channelIdentifierMap != null) { Class[] interfaces = type.getInterfaces(); // first try to find a match amongst the interfaces and also check if there is more then one for (Class interfase : interfaces) { diff --git a/spring-integration-core/src/test/java/org/springframework/integration/core/MessagingTemplateTests.java b/spring-integration-core/src/test/java/org/springframework/integration/core/MessagingTemplateTests.java index d5e120b29c..e3d48e8192 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/core/MessagingTemplateTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/core/MessagingTemplateTests.java @@ -47,6 +47,7 @@ import org.springframework.integration.scheduling.PollerMetadata; import org.springframework.integration.support.MessageBuilder; import org.springframework.integration.support.channel.BeanFactoryChannelResolver; import org.springframework.integration.support.channel.ChannelResolutionException; +import org.springframework.integration.support.channel.ChannelResolver; import org.springframework.integration.support.converter.SimpleMessageConverter; import org.springframework.integration.test.util.TestUtils; import org.springframework.integration.test.util.TestUtils.TestApplicationContext; @@ -307,16 +308,27 @@ public class MessagingTemplateTests { @Test public void sendByChannelNameWithCustomChannelResolver() { QueueChannel testChannel = new QueueChannel(); + final QueueChannel anotherChannel = new QueueChannel(); DefaultListableBeanFactory beanFactory = new DefaultListableBeanFactory(); beanFactory.registerSingleton("testChannel", testChannel); MessagingTemplate template = new MessagingTemplate(); template.setBeanFactory(beanFactory); + template.afterPropertiesSet(); Message message = MessageBuilder.withPayload("test").build(); template.send("testChannel", message); assertEquals(message, testChannel.receive(0)); + + template.setChannelResolver(new ChannelResolver() { + public MessageChannel resolveChannelName(String channelName) { + return anotherChannel; + } + }); + message = MessageBuilder.withPayload("test").build(); + template.send("testChannel", message); + assertEquals(message, anotherChannel.receive(0)); } @Test(expected = IllegalStateException.class) diff --git a/spring-integration-core/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java b/spring-integration-core/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java index 9989315fae..0ee38d90f1 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java @@ -16,6 +16,7 @@ package org.springframework.integration.router; +import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertSame; @@ -25,10 +26,12 @@ import org.springframework.beans.factory.support.ManagedMap; import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.context.support.StaticApplicationContext; import org.springframework.integration.Message; +import org.springframework.integration.MessageChannel; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.core.MessageHandler; import org.springframework.integration.support.MessageBuilder; import org.springframework.integration.support.channel.BeanFactoryChannelResolver; +import org.springframework.integration.support.channel.ChannelResolver; /** * @author Mark Fisher @@ -61,6 +64,7 @@ public class HeaderValueRouterTests { routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); context.registerBeanDefinition("router", routerBeanDefinition); context.registerBeanDefinition("testChannel", new RootBeanDefinition(QueueChannel.class)); + context.registerBeanDefinition("newChannel", new RootBeanDefinition(QueueChannel.class)); context.refresh(); MessageHandler handler = (MessageHandler) context.getBean("router"); Message message = MessageBuilder.withPayload("test").setHeader("testHeaderName", "testChannel").build(); @@ -69,6 +73,20 @@ public class HeaderValueRouterTests { Message result = channel.receive(1000); assertNotNull(result); assertSame(message, result); + + // validate dynamics + HeaderValueRouter router = (HeaderValueRouter) context.getBean("router"); + router.setChannelMapping("testChannel", "newChannel"); + router.handleMessage(message); + QueueChannel newChannel = (QueueChannel) context.getBean("newChannel"); + result = newChannel.receive(10); + assertNotNull(result); + + router.removeChannelMapping("testChannel"); + router.handleMessage(message); + result = channel.receive(1000); + assertNotNull(result); + assertSame(message, result); } @Test @@ -77,14 +95,11 @@ public class HeaderValueRouterTests { StaticApplicationContext context = new StaticApplicationContext(); ManagedMap channelMap = new ManagedMap(); channelMap.put("testKey", "testChannel"); - RootBeanDefinition channelResolverBeanDefinition = new RootBeanDefinition(BeanFactoryChannelResolver.class); - channelResolverBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue(context); RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); routerBeanDefinition.getPropertyValues().addPropertyValue("channelIdentifierMap", channelMap); - routerBeanDefinition.getPropertyValues().addPropertyValue("channelResolver", new RuntimeBeanReference("resolver")); - context.registerBeanDefinition("resolver", channelResolverBeanDefinition); + routerBeanDefinition.getPropertyValues().addPropertyValue("beanFactory", context); context.registerBeanDefinition("router", routerBeanDefinition); context.registerBeanDefinition("testChannel", new RootBeanDefinition(QueueChannel.class)); context.refresh(); @@ -96,6 +111,34 @@ public class HeaderValueRouterTests { assertNotNull(result); assertSame(message, result); } + @Test + @SuppressWarnings("unchecked") + public void resolveChannelNameFromMapAndCustomeResolver() { + final StaticApplicationContext context = new StaticApplicationContext(); + ManagedMap channelMap = new ManagedMap(); + channelMap.put("testKey", "testChannel"); + RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); + routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); + routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); + routerBeanDefinition.getPropertyValues().addPropertyValue("channelIdentifierMap", channelMap); + routerBeanDefinition.getPropertyValues().addPropertyValue("beanFactory", context); + routerBeanDefinition.getPropertyValues().addPropertyValue("channelResolver", new ChannelResolver() { + public MessageChannel resolveChannelName(String channelName) { + return context.getBean("anotherChannel", MessageChannel.class); + } + }); + context.registerBeanDefinition("router", routerBeanDefinition); + context.registerBeanDefinition("testChannel", new RootBeanDefinition(QueueChannel.class)); + context.registerBeanDefinition("anotherChannel", new RootBeanDefinition(QueueChannel.class)); + context.refresh(); + MessageHandler handler = (MessageHandler) context.getBean("router"); + Message message = MessageBuilder.withPayload("test").setHeader("testHeaderName", "testKey").build(); + handler.handleMessage(message); + QueueChannel channel = (QueueChannel) context.getBean("anotherChannel"); + Message result = channel.receive(1000); + assertNotNull(result); + assertSame(message, result); + } @Test public void resolveMultipleChannelsWithStringArray() { @@ -145,4 +188,5 @@ public class HeaderValueRouterTests { assertSame(message, result2); } + } diff --git a/spring-integration-core/src/test/java/org/springframework/integration/router/PayloadTypeRouterTests.java b/spring-integration-core/src/test/java/org/springframework/integration/router/PayloadTypeRouterTests.java index ad2a628b60..f38c8478ae 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/router/PayloadTypeRouterTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/router/PayloadTypeRouterTests.java @@ -16,6 +16,7 @@ package org.springframework.integration.router; +import static junit.framework.Assert.fail; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertNull; @@ -31,7 +32,6 @@ import org.springframework.integration.MessageChannel; import org.springframework.integration.MessageHandlingException; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.message.GenericMessage; -import org.springframework.integration.support.channel.BeanFactoryChannelResolver; /** * @author Mark Fisher @@ -57,11 +57,32 @@ public class PayloadTypeRouterTests { Message message1 = new GenericMessage("test"); Message message2 = new GenericMessage(123); + assertEquals(1, router.determineTargetChannels(message1).size()); MessageChannel result1 = router.determineTargetChannels(message1).iterator().next(); + assertEquals(1, router.determineTargetChannels(message2).size()); MessageChannel result2 = router.determineTargetChannels(message2).iterator().next(); assertEquals(stringChannel, result1); assertEquals(integerChannel, result2); + // validate dynamics + QueueChannel newChannel = new QueueChannel(); + beanFactory.registerSingleton("newChannel", newChannel); + router.setChannelMapping(String.class.getName(), "newChannel"); + assertEquals(1, router.determineTargetChannels(message1).size()); + result1 = router.determineTargetChannels(message1).iterator().next(); + assertEquals(newChannel, result1); + // validate nothing happens if mappings were removed and resolutionRequires = false + router.removeChannelMapping(String.class.getName()); + router.removeChannelMapping(Integer.class.getName()); + router.handleMessage(message1); + // validate exception is thrown if mappings were removed and resolutionRequires = true + router.setResolutionRequired(true); + try { + router.handleMessage(message1); + fail(); + } catch (Exception e) { + // ignore + } } @Test @@ -86,6 +107,15 @@ public class PayloadTypeRouterTests { assertNotNull(result); assertEquals(99, result.getPayload()); assertNull(defaultChannel.receive(0)); + + // validate dynamics + QueueChannel newChannel = new QueueChannel(); + beanFactory.registerSingleton("newChannel", newChannel); + router.setChannelMapping(Integer.class.getName(), "newChannel"); + assertEquals(1, router.determineTargetChannels(message).size()); + router.handleMessage(message); + result = newChannel.receive(10); + assertNotNull(result); } @Test @@ -177,6 +207,15 @@ public class PayloadTypeRouterTests { assertEquals(99, result.getPayload()); assertNull(numberChannel.receive(0)); assertNull(defaultChannel.receive(0)); + + // validate dynamics + QueueChannel newChannel = new QueueChannel(); + beanFactory.registerSingleton("newChannel", newChannel); + router.setChannelMapping(Integer.class.getName(), "newChannel"); + assertEquals(1, router.determineTargetChannels(message).size()); + router.handleMessage(message); + result = newChannel.receive(10); + assertNotNull(result); } @Test(expected = IllegalStateException.class) @@ -305,5 +344,4 @@ public class PayloadTypeRouterTests { assertNotNull(result2); assertEquals(123, result2.getPayload()); } - } diff --git a/spring-integration-core/src/test/java/org/springframework/integration/router/RouterTests.java b/spring-integration-core/src/test/java/org/springframework/integration/router/RouterTests.java index 53eaef25d3..eb71440f2b 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/router/RouterTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/router/RouterTests.java @@ -20,25 +20,22 @@ import static org.junit.Assert.assertEquals; import static org.mockito.Mockito.mock; import java.util.ArrayList; -import java.util.Collections; import java.util.List; import org.junit.Test; - import org.springframework.beans.factory.BeanFactory; import org.springframework.context.support.GenericApplicationContext; import org.springframework.integration.Message; -import org.springframework.integration.MessageChannel; import org.springframework.integration.MessageDeliveryException; import org.springframework.integration.MessagingException; import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.channel.TestChannelResolver; import org.springframework.integration.message.GenericMessage; -import org.springframework.integration.support.channel.BeanFactoryChannelResolver; import org.springframework.util.CollectionUtils; /** * @author Mark Fisher + * @author Oleg Zhurakousky */ public class RouterTests { @@ -48,8 +45,7 @@ public class RouterTests { @Override protected List getChannelIndicatorList(Message message) { return null; - } - + } }; Message message = new GenericMessage("test"); router.handleMessage(message);