diff --git a/org.springframework.integration/src/main/java/org/springframework/integration/config/xml/ChainParser.java b/org.springframework.integration/src/main/java/org/springframework/integration/config/xml/ChainParser.java index c663647ce3..2fe05c16fa 100644 --- a/org.springframework.integration/src/main/java/org/springframework/integration/config/xml/ChainParser.java +++ b/org.springframework.integration/src/main/java/org/springframework/integration/config/xml/ChainParser.java @@ -27,19 +27,21 @@ import org.springframework.beans.factory.support.BeanDefinitionBuilder; import org.springframework.beans.factory.support.BeanDefinitionReaderUtils; import org.springframework.beans.factory.support.ManagedList; import org.springframework.beans.factory.xml.ParserContext; +import org.springframework.integration.handler.AbstractReplyProducingMessageHandler; /** * Parser for the <chain> element. * * @author Mark Fisher + * @author Iwein Fuld */ public class ChainParser extends AbstractConsumerEndpointParser { @Override @SuppressWarnings("unchecked") protected BeanDefinitionBuilder parseHandler(Element element, ParserContext parserContext) { - BeanDefinitionBuilder builder = BeanDefinitionBuilder.genericBeanDefinition( - IntegrationNamespaceUtils.BASE_PACKAGE + ".handler.MessageHandlerChain"); + BeanDefinitionBuilder builder = BeanDefinitionBuilder + .genericBeanDefinition(IntegrationNamespaceUtils.BASE_PACKAGE + ".handler.MessageHandlerChain"); ManagedList handlerList = new ManagedList(); NodeList children = element.getChildNodes(); for (int i = 0; i < children.getLength(); i++) { @@ -54,7 +56,13 @@ public class ChainParser extends AbstractConsumerEndpointParser { } private String parseChild(Element element, ParserContext parserContext, BeanDefinition parentDefinition) { - BeanDefinition beanDefinition = parserContext.getDelegate().parseCustomElement(element, parentDefinition); + BeanDefinition beanDefinition; + if (element.getLocalName().equals("bean")) { + beanDefinition = parserContext.getDelegate().parseBeanDefinitionElement(element).getBeanDefinition(); + } + else { + beanDefinition = parserContext.getDelegate().parseCustomElement(element, parentDefinition); + } if (beanDefinition == null) { parserContext.getReaderContext().error("child BeanDefinition must not be null", element); } diff --git a/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests-context.xml b/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests-context.xml index 91425a9643..274a91ab30 100644 --- a/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests-context.xml +++ b/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests-context.xml @@ -1,58 +1,62 @@ - - - - + - + - + - - + + - -
-
+ +
+
- + - + - - + + + + + + - + - - + + - - - - - + + + + diff --git a/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests.java b/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests.java index 699da0e4a4..d01b7a2343 100644 --- a/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests.java +++ b/org.springframework.integration/src/test/java/org/springframework/integration/config/ChainParserTests.java @@ -16,26 +16,35 @@ package org.springframework.integration.config; +import static org.hamcrest.CoreMatchers.is; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertNull; +import static org.junit.Assert.assertThat; import org.junit.Test; +import org.junit.runner.RunWith; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.beans.factory.annotation.Qualifier; import org.springframework.integration.channel.PollableChannel; import org.springframework.integration.core.Message; import org.springframework.integration.core.MessageChannel; +import org.springframework.integration.handler.AbstractReplyProducingMessageHandler; +import org.springframework.integration.handler.ReplyMessageHolder; import org.springframework.integration.message.MessageBuilder; +import org.springframework.integration.message.MessageHandler; import org.springframework.test.context.ContextConfiguration; import org.springframework.test.context.junit4.AbstractJUnit4SpringContextTests; +import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; /** * @author Mark Fisher + * @author Iwein Fuld */ @ContextConfiguration -public class ChainParserTests extends AbstractJUnit4SpringContextTests { +@RunWith(SpringJUnit4ClassRunner.class) +public class ChainParserTests { @Autowired @Qualifier("filterInput") @@ -57,6 +66,13 @@ public class ChainParserTests extends AbstractJUnit4SpringContextTests { @Qualifier("replyOutput") private PollableChannel replyOutput; + + @Autowired + @Qualifier("beanInput") + private MessageChannel beanInput; + + public static Message successMessage = MessageBuilder.withPayload("success").build(); + @Test public void chainWithAcceptingFilter() { @@ -95,5 +111,22 @@ public class ChainParserTests extends AbstractJUnit4SpringContextTests { assertNotNull(reply); assertEquals("foo", reply.getPayload()); } + + @Test + public void chainHandlerBean() throws Exception { + Message message = MessageBuilder.withPayload("test").build(); + this.beanInput.send(message); + Message reply = this.output.receive(3000); + assertNotNull(reply); + assertThat(reply, is(successMessage)); + } + + public static class StubHandler extends AbstractReplyProducingMessageHandler { + @Override + protected void handleRequestMessage(Message requestMessage, ReplyMessageHolder replyMessageHolder) { + replyMessageHolder.add(successMessage); + } + + } }