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 new file mode 100644 index 0000000000..752b127835 --- /dev/null +++ b/spring-integration-core/src/main/java/org/springframework/integration/router/PayloadTypeRouter.java @@ -0,0 +1,62 @@ +/* + * Copyright 2002-2007 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.integration.router; + +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + +import org.springframework.integration.channel.MessageChannel; +import org.springframework.integration.message.Message; +import org.springframework.util.Assert; + +/** + * A router implementation that resolves the {@link MessageChannel} based on the + * {@link Message Message's} payload type. + * + * @author Mark Fisher + */ +public class PayloadTypeRouter extends SingleChannelRouter { + + private Map, MessageChannel> channelMappings = new ConcurrentHashMap, MessageChannel>(); + + private MessageChannel defaultChannel; + + + public PayloadTypeRouter() { + this.setChannelResolver(new PayloadTypeChannelResolver()); + } + + + public void setChannelMappings(Map, MessageChannel> channelMappings) { + Assert.notNull(channelMappings, "'channelMappings' must not be null"); + this.channelMappings = channelMappings; + } + + public void setDefaultChannel(MessageChannel defaultChannel) { + this.defaultChannel = defaultChannel; + } + + + private class PayloadTypeChannelResolver implements ChannelResolver { + + public MessageChannel resolve(Message message) { + MessageChannel channel = channelMappings.get(message.getPayload().getClass()); + return channel != null ? channel : defaultChannel; + } + } + +} 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 new file mode 100644 index 0000000000..1ea88de8f6 --- /dev/null +++ b/spring-integration-core/src/test/java/org/springframework/integration/router/PayloadTypeRouterTests.java @@ -0,0 +1,82 @@ +/* + * Copyright 2002-2007 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.integration.router; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotNull; + +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + +import org.junit.Test; + +import org.springframework.integration.channel.MessageChannel; +import org.springframework.integration.channel.PointToPointChannel; +import org.springframework.integration.message.GenericMessage; +import org.springframework.integration.message.Message; +import org.springframework.integration.message.StringMessage; + +/** + * @author Mark Fisher + */ +public class PayloadTypeRouterTests { + + @Test + public void testRoutingByPayloadType() { + PointToPointChannel stringChannel = new PointToPointChannel(); + PointToPointChannel integerChannel = new PointToPointChannel(); + Map, MessageChannel> channelMappings = new ConcurrentHashMap, MessageChannel>(); + channelMappings.put(String.class, stringChannel); + channelMappings.put(Integer.class, integerChannel); + PayloadTypeRouter router = new PayloadTypeRouter(); + router.setChannelMappings(channelMappings); + router.afterPropertiesSet(); + Message message1 = new StringMessage("test"); + Message message2 = new GenericMessage(123); + router.handle(message1); + router.handle(message2); + Message result1 = stringChannel.receive(25); + assertNotNull(result1); + assertEquals("test", result1.getPayload()); + Message result2 = integerChannel.receive(25); + assertNotNull(result2); + assertEquals(123, result2.getPayload()); + } + + @Test + public void testRoutingToDefaultChannelWhenNoTypeMatches() { + PointToPointChannel stringChannel = new PointToPointChannel(); + PointToPointChannel defaultChannel = new PointToPointChannel(); + Map, MessageChannel> channelMappings = new ConcurrentHashMap, MessageChannel>(); + channelMappings.put(String.class, stringChannel); + PayloadTypeRouter router = new PayloadTypeRouter(); + router.setChannelMappings(channelMappings); + router.setDefaultChannel(defaultChannel); + router.afterPropertiesSet(); + Message message1 = new StringMessage("test"); + Message message2 = new GenericMessage(123); + router.handle(message1); + router.handle(message2); + Message result1 = stringChannel.receive(25); + assertNotNull(result1); + assertEquals("test", result1.getPayload()); + Message result2 = defaultChannel.receive(25); + assertNotNull(result2); + assertEquals(123, result2.getPayload()); + } + +}