diff --git a/org.springframework.integration/src/main/java/org/springframework/integration/channel/MapBasedChannelResolver.java b/org.springframework.integration/src/main/java/org/springframework/integration/channel/MapBasedChannelResolver.java index 6db3672139..20e912f1d4 100644 --- a/org.springframework.integration/src/main/java/org/springframework/integration/channel/MapBasedChannelResolver.java +++ b/org.springframework.integration/src/main/java/org/springframework/integration/channel/MapBasedChannelResolver.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2008 the original author or authors. + * Copyright 2002-2009 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. @@ -32,7 +32,25 @@ public class MapBasedChannelResolver implements ChannelResolver { private volatile Map channelMap = new HashMap(); + /** + * Empty constructor for use when providing the channel map via + * {@link #setChannelMap(Map)}. + */ + public MapBasedChannelResolver() { + } + /** + * Create a {@link ChannelResolver} that uses the provided Map. + * Each String key will resolve to the associated channel value. + */ + public MapBasedChannelResolver(Map channelMap) { + this.setChannelMap(channelMap); + } + + /** + * Provide a map of channels to be used by this resolver. + * Each String key will resolve to the associated channel value. + */ public void setChannelMap(Map channelMap) { Assert.notNull(channelMap, "channelMap must not be null"); this.channelMap = channelMap; diff --git a/org.springframework.integration/src/main/java/org/springframework/integration/router/HeaderValueRouter.java b/org.springframework.integration/src/main/java/org/springframework/integration/router/HeaderValueRouter.java new file mode 100644 index 0000000000..897ae487ee --- /dev/null +++ b/org.springframework.integration/src/main/java/org/springframework/integration/router/HeaderValueRouter.java @@ -0,0 +1,54 @@ +/* + * Copyright 2002-2009 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.Collections; +import java.util.List; + +import org.springframework.integration.core.Message; +import org.springframework.util.Assert; +import org.springframework.util.StringUtils; + +/** + * A Message Router that resolves the MessageChannel from a header value. + * + * @author Oleg Zhurakousky + * @author Mark Fisher + * @Since 1.0.3 + */ +public class HeaderValueRouter extends AbstractChannelNameResolvingMessageRouter { + + private final String headerName; + + /** + * Create a router that uses the provided header name to lookup a channel. + */ + public HeaderValueRouter(String headerName) { + Assert.notNull(headerName, "'headerName' must not be null"); + this.headerName = headerName; + } + + @Override + protected List getChannelIndicatorList(Message message) { + Object value = message.getHeaders().get(this.headerName); + if (value instanceof String && ((String) value).indexOf(',') != -1) { + value = StringUtils.tokenizeToStringArray((String) value, ",", true, true); + } + return Collections.singletonList(value); + } + +} diff --git a/org.springframework.integration/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java b/org.springframework.integration/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java new file mode 100644 index 0000000000..cad98db68d --- /dev/null +++ b/org.springframework.integration/src/test/java/org/springframework/integration/router/HeaderValueRouterTests.java @@ -0,0 +1,147 @@ +/* + * Copyright 2002-2009 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.assertNotNull; +import static org.junit.Assert.assertSame; + +import org.junit.Test; + +import org.springframework.beans.factory.config.RuntimeBeanReference; +import org.springframework.beans.factory.support.ManagedMap; +import org.springframework.beans.factory.support.RootBeanDefinition; +import org.springframework.context.support.StaticApplicationContext; +import org.springframework.integration.channel.MapBasedChannelResolver; +import org.springframework.integration.channel.QueueChannel; +import org.springframework.integration.core.Message; +import org.springframework.integration.message.MessageBuilder; +import org.springframework.integration.message.MessageHandler; + +/** + * @author Mark Fisher + */ +public class HeaderValueRouterTests { + + @Test + public void channelAsHeaderValue() { + StaticApplicationContext context = new StaticApplicationContext(); + RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); + routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); + routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); + context.registerBeanDefinition("router", routerBeanDefinition); + context.refresh(); + MessageHandler handler = (MessageHandler) context.getBean("router"); + QueueChannel testChannel = new QueueChannel(); + Message message = MessageBuilder.withPayload("test").setHeader("testHeaderName", testChannel).build(); + handler.handleMessage(message); + Message result = testChannel.receive(1000); + assertNotNull(result); + assertSame(message, result); + } + + @Test + public void resolveChannelNameFromContext() { + StaticApplicationContext context = new StaticApplicationContext(); + RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); + routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); + routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); + context.registerBeanDefinition("router", routerBeanDefinition); + context.registerBeanDefinition("testChannel", new RootBeanDefinition(QueueChannel.class)); + context.refresh(); + MessageHandler handler = (MessageHandler) context.getBean("router"); + Message message = MessageBuilder.withPayload("test").setHeader("testHeaderName", "testChannel").build(); + handler.handleMessage(message); + QueueChannel channel = (QueueChannel) context.getBean("testChannel"); + Message result = channel.receive(1000); + assertNotNull(result); + assertSame(message, result); + } + + @Test + @SuppressWarnings("unchecked") + public void resolveChannelNameFromMap() { + StaticApplicationContext context = new StaticApplicationContext(); + ManagedMap channelMap = new ManagedMap(); + channelMap.put("testKey", new RuntimeBeanReference("testChannel")); + RootBeanDefinition channelResolverBeanDefinition = new RootBeanDefinition(MapBasedChannelResolver.class); + channelResolverBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue(channelMap); + RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); + routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); + routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); + routerBeanDefinition.getPropertyValues().addPropertyValue("channelResolver", new RuntimeBeanReference("resolver")); + context.registerBeanDefinition("resolver", channelResolverBeanDefinition); + context.registerBeanDefinition("router", routerBeanDefinition); + context.registerBeanDefinition("testChannel", 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("testChannel"); + Message result = channel.receive(1000); + assertNotNull(result); + assertSame(message, result); + } + + @Test + public void resolveMultipleChannelsWithStringArray() { + StaticApplicationContext context = new StaticApplicationContext(); + RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); + routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); + routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); + context.registerBeanDefinition("router", routerBeanDefinition); + context.registerBeanDefinition("channel1", new RootBeanDefinition(QueueChannel.class)); + context.registerBeanDefinition("channel2", new RootBeanDefinition(QueueChannel.class)); + context.refresh(); + MessageHandler handler = (MessageHandler) context.getBean("router"); + String[] channels = new String[] { "channel1", "channel2" }; + Message message = MessageBuilder.withPayload("test").setHeader("testHeaderName", channels).build(); + handler.handleMessage(message); + QueueChannel channel1 = (QueueChannel) context.getBean("channel1"); + QueueChannel channel2 = (QueueChannel) context.getBean("channel2"); + Message result1 = channel1.receive(1000); + Message result2 = channel2.receive(1000); + assertNotNull(result1); + assertNotNull(result2); + assertSame(message, result1); + assertSame(message, result2); + } + + @Test + public void resolveMultipleChannelsWithCommaDelimitedString() { + StaticApplicationContext context = new StaticApplicationContext(); + RootBeanDefinition routerBeanDefinition = new RootBeanDefinition(HeaderValueRouter.class); + routerBeanDefinition.getConstructorArgumentValues().addGenericArgumentValue("testHeaderName"); + routerBeanDefinition.getPropertyValues().addPropertyValue("resolutionRequired", "true"); + context.registerBeanDefinition("router", routerBeanDefinition); + context.registerBeanDefinition("channel1", new RootBeanDefinition(QueueChannel.class)); + context.registerBeanDefinition("channel2", new RootBeanDefinition(QueueChannel.class)); + context.refresh(); + MessageHandler handler = (MessageHandler) context.getBean("router"); + String channels = "channel1, channel2"; + Message message = MessageBuilder.withPayload("test").setHeader("testHeaderName", channels).build(); + handler.handleMessage(message); + QueueChannel channel1 = (QueueChannel) context.getBean("channel1"); + QueueChannel channel2 = (QueueChannel) context.getBean("channel2"); + Message result1 = channel1.receive(1000); + Message result2 = channel2.receive(1000); + assertNotNull(result1); + assertNotNull(result2); + assertSame(message, result1); + assertSame(message, result2); + } + +}