From 1f0d9d6e45276a15cc6ec925c7ce41a76476cb88 Mon Sep 17 00:00:00 2001 From: Mark Fisher Date: Tue, 17 Aug 2010 19:15:27 +0000 Subject: [PATCH] INT-1354 added support for "expression" attributes on the
sub-elements of a --- .../integration/config/xml/GatewayParser.java | 30 ++++++- .../gateway/GatewayMethodDefinition.java | 12 +-- .../gateway/GatewayProxyFactoryBean.java | 8 +- .../handler/ArgumentArrayMessageMapper.java | 35 ++++---- .../HeaderEnrichedGatewayTests-context.xml | 37 ++++++-- .../gateway/HeaderEnrichedGatewayTests.java | 90 +++++++++++-------- ...umentArrayMessageMapperToMessageTests.java | 15 ++-- 7 files changed, 149 insertions(+), 78 deletions(-) diff --git a/spring-integration-core/src/main/java/org/springframework/integration/config/xml/GatewayParser.java b/spring-integration-core/src/main/java/org/springframework/integration/config/xml/GatewayParser.java index 0a1f855557..f2e48873b4 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/config/xml/GatewayParser.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/config/xml/GatewayParser.java @@ -21,12 +21,15 @@ import java.util.Map; import org.w3c.dom.Element; +import org.springframework.beans.factory.BeanDefinitionStoreException; import org.springframework.beans.factory.config.BeanDefinition; import org.springframework.beans.factory.support.BeanDefinitionBuilder; import org.springframework.beans.factory.support.ManagedMap; +import org.springframework.beans.factory.support.RootBeanDefinition; import org.springframework.beans.factory.xml.AbstractSimpleBeanDefinitionParser; import org.springframework.util.CollectionUtils; import org.springframework.util.ObjectUtils; +import org.springframework.util.StringUtils; import org.springframework.util.xml.DomUtils; /** @@ -103,12 +106,31 @@ public class GatewayParser extends AbstractSimpleBeanDefinitionParser { builder.addPropertyValue("methodToChannelMap", methodToChannelMap); } - private void setMethodInvocationHeaders(BeanDefinitionBuilder gatewayDefinitionBuilder, List invocationHeaders){ - Map methodInvocationHeaders = new ManagedMap(); + private void setMethodInvocationHeaders(BeanDefinitionBuilder gatewayDefinitionBuilder, List invocationHeaders) { + Map headerExpressions = new ManagedMap(); for (Element headerElement : invocationHeaders) { - methodInvocationHeaders.put(headerElement.getAttribute("name"), headerElement.getAttribute("value")); + String headerName = headerElement.getAttribute("name"); + String headerValue = headerElement.getAttribute("value"); + String headerExpression = headerElement.getAttribute("expression"); + boolean hasValue = StringUtils.hasText(headerValue); + boolean hasExpression = StringUtils.hasText(headerExpression); + if (!(hasValue ^ hasExpression)) { + throw new BeanDefinitionStoreException("exactly one of 'value' or 'expression' is required on a header sub-element"); + } + RootBeanDefinition expressionDef = null; + if (hasValue) { + expressionDef = new RootBeanDefinition("org.springframework.expression.common.LiteralExpression"); + expressionDef.getConstructorArgumentValues().addGenericArgumentValue(headerValue); + } + else if (hasExpression) { + expressionDef = new RootBeanDefinition("org.springframework.integration.config.ExpressionFactoryBean"); + expressionDef.getConstructorArgumentValues().addGenericArgumentValue(headerExpression); + } + if (expressionDef != null) { + headerExpressions.put(headerName, expressionDef); + } } - gatewayDefinitionBuilder.addPropertyValue("staticHeaders", methodInvocationHeaders); + gatewayDefinitionBuilder.addPropertyValue("headerExpressions", headerExpressions); } } diff --git a/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayMethodDefinition.java b/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayMethodDefinition.java index f2c44b3d27..508db86b3d 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayMethodDefinition.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayMethodDefinition.java @@ -19,6 +19,8 @@ package org.springframework.integration.gateway; import java.util.HashMap; import java.util.Map; +import org.springframework.expression.Expression; + /** * Represents the definition of Gateway methods, when using multiple methods per Gateway interface. * <si:method name="echo" request-channel="inputA" reply-timeout="2" request-timeout="200"/> @@ -38,7 +40,7 @@ public class GatewayMethodDefinition { private volatile String replyTimeout; - private volatile Map staticHeaders = new HashMap(); + private volatile Map headerExpressions = new HashMap(); public String getPayloadExpression() { @@ -49,12 +51,12 @@ public class GatewayMethodDefinition { this.payloadExpression = payloadExpression; } - public Map getStaticHeaders() { - return staticHeaders; + public Map getHeaderExpressions() { + return this.headerExpressions; } - public void setStaticHeaders(Map staticHeaders) { - this.staticHeaders = staticHeaders; + public void setHeaderExpressions(Map headerExpressions) { + this.headerExpressions = headerExpressions; } public String getRequestChannelName() { diff --git a/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayProxyFactoryBean.java b/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayProxyFactoryBean.java index 0dd122a509..2d5f8e7d98 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayProxyFactoryBean.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/gateway/GatewayProxyFactoryBean.java @@ -24,6 +24,7 @@ import java.util.Map; import org.aopalliance.intercept.MethodInterceptor; import org.aopalliance.intercept.MethodInvocation; + import org.springframework.aop.framework.ProxyFactory; import org.springframework.aop.support.AopUtils; import org.springframework.beans.SimpleTypeConverter; @@ -32,6 +33,7 @@ import org.springframework.beans.factory.BeanClassLoaderAware; import org.springframework.beans.factory.BeanFactory; import org.springframework.beans.factory.FactoryBean; import org.springframework.core.convert.ConversionService; +import org.springframework.expression.Expression; import org.springframework.integration.Message; import org.springframework.integration.annotation.Gateway; import org.springframework.integration.context.BeanFactoryChannelResolver; @@ -275,7 +277,7 @@ public class GatewayProxyFactoryBean extends AbstractEndpoint implements Factory long requestTimeout = this.defaultRequestTimeout; long replyTimeout = this.defaultReplyTimeout; String payloadExpression = null; - Map staticHeaders = null; + Map headerExpressions = null; if (gatewayAnnotation != null) { String requestChannelName = gatewayAnnotation.requestChannel(); if (StringUtils.hasText(requestChannelName)) { @@ -292,7 +294,7 @@ public class GatewayProxyFactoryBean extends AbstractEndpoint implements Factory GatewayMethodDefinition gatewayDefinition = methodToChannelMap.get(method.getName()); if (gatewayDefinition != null) { payloadExpression = gatewayDefinition.getPayloadExpression(); - staticHeaders = gatewayDefinition.getStaticHeaders(); + headerExpressions = gatewayDefinition.getHeaderExpressions(); String requestChannelName = gatewayDefinition.getRequestChannelName(); if (StringUtils.hasText(requestChannelName)) { requestChannel = this.resolveChannelName(requestChannelName); @@ -311,7 +313,7 @@ public class GatewayProxyFactoryBean extends AbstractEndpoint implements Factory } } } - ArgumentArrayMessageMapper messageMapper = new ArgumentArrayMessageMapper(method, staticHeaders); + ArgumentArrayMessageMapper messageMapper = new ArgumentArrayMessageMapper(method, headerExpressions); if (StringUtils.hasText(payloadExpression)) { messageMapper.setPayloadExpression(payloadExpression); } diff --git a/spring-integration-core/src/main/java/org/springframework/integration/handler/ArgumentArrayMessageMapper.java b/spring-integration-core/src/main/java/org/springframework/integration/handler/ArgumentArrayMessageMapper.java index 71aa236f3c..a3790daae6 100644 --- a/spring-integration-core/src/main/java/org/springframework/integration/handler/ArgumentArrayMessageMapper.java +++ b/spring-integration-core/src/main/java/org/springframework/integration/handler/ArgumentArrayMessageMapper.java @@ -108,7 +108,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper staticHeaders; + private final Map headerExpressions; private final List parameterList; @@ -116,7 +116,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper parameterPayloadExpressions = new HashMap(); - private final StandardEvaluationContext staticEvaluationContext = new StandardEvaluationContext(); + private final StandardEvaluationContext evaluationContext = new StandardEvaluationContext(); private volatile BeanResolver beanResolver; @@ -125,10 +125,10 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper staticHeaders) { + public ArgumentArrayMessageMapper(Method method, Map headerExpressions) { Assert.notNull(method, "method must not be null"); this.method = method; - this.staticHeaders = staticHeaders; + this.headerExpressions = headerExpressions; this.parameterList = getMethodParameterList(method); this.payloadExpression = parsePayloadExpression(method); } @@ -141,7 +141,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper mapArgumentsToMessage(Object[] arguments) { Object messageOrPayload = null; boolean foundPayloadAnnotation = false; @@ -199,10 +198,10 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper) argumentValue).keySet()) { Assert.isInstanceOf(String.class, key, "Invalid header name [" + key + "], name type must be String."); - Object value = ((Map) argumentValue).get(key); + Object value = ((Map) argumentValue).get(key); headers.put((String) key, value); } } @@ -216,19 +215,26 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper) argumentValue, headers); } else if (this.payloadExpression == null) { this.throwExceptionForMultipleMessageOrPayloadParameters(methodParameter); } } Assert.isTrue(messageOrPayload != null, "unable to determine a Message or payload parameter on method [" + method + "]"); - MessageBuilder builder = (messageOrPayload instanceof Message) + MessageBuilder builder = (messageOrPayload instanceof Message) ? MessageBuilder.fromMessage((Message) messageOrPayload) : MessageBuilder.withPayload(messageOrPayload); builder.copyHeadersIfAbsent(headers); - if (!CollectionUtils.isEmpty(staticHeaders)){ - builder.copyHeaders(staticHeaders); + if (!CollectionUtils.isEmpty(this.headerExpressions)) { + Map evaluatedHeaders = new HashMap(); + for (Map.Entry entry : this.headerExpressions.entrySet()) { + Object value = entry.getValue().getValue(this.evaluationContext); + if (value != null) { + evaluatedHeaders.put(entry.getKey(), value); + } + } + builder.copyHeaders(evaluatedHeaders); } return builder.build(); } @@ -239,7 +245,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper headers) { + private void copyHeaders(Map argumentValue, Map headers) { for (Object key : argumentValue.keySet()) { if (!(key instanceof String)) { throw new IllegalArgumentException("Invalid header name [" + key + diff --git a/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests-context.xml b/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests-context.xml index 8a6c157a2a..1e0270d331 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests-context.xml +++ b/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests-context.xml @@ -5,25 +5,48 @@ http://www.springframework.org/schema/integration http://www.springframework.org/schema/integration/spring-integration-2.0.xsd" xmlns:int="http://www.springframework.org/schema/integration"> - - - + + - + - + + + + + + + + + + + + + + + + - + + + + + - + + + + diff --git a/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests.java b/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests.java index a0872f3b8e..36a269f129 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/gateway/HeaderEnrichedGatewayTests.java @@ -21,21 +21,18 @@ import static junit.framework.Assert.assertNull; import org.junit.Test; import org.junit.runner.RunWith; -import org.mockito.Mockito; -import org.mockito.invocation.InvocationOnMock; -import org.mockito.stubbing.Answer; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.integration.Message; import org.springframework.integration.MessageHeaders; import org.springframework.integration.annotation.Header; -import org.springframework.integration.channel.DirectChannel; -import org.springframework.integration.core.MessageHandler; +import org.springframework.integration.core.PollableChannel; import org.springframework.test.context.ContextConfiguration; import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; /** * @author Oleg Zhurakousky + * @author Mark Fisher * @since 2.0 */ @RunWith(SpringJUnit4ClassRunner.class) @@ -43,52 +40,69 @@ import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; public class HeaderEnrichedGatewayTests { @Autowired - private SampleGateway gateway; + private SampleGateway gatewayWithHeaderValues; @Autowired - private DirectChannel input; + private SampleGateway gatewayWithHeaderExpressions; + + @Autowired + private PollableChannel channel; private Object testPayload; @Test - public void validateStaticHeaderMappings() throws Exception { - MessageHandler handler = Mockito.mock(MessageHandler.class); - input.subscribe(handler); - - this.prepareHandlerForTest(handler); + public void validateHeaderValueMappings() throws Exception { testPayload = "hello"; - gateway.sendString((String) testPayload); - Mockito.verify(handler, Mockito.times(1)).handleMessage(Mockito.any(Message.class)); - - this.prepareHandlerForTest(handler); + gatewayWithHeaderValues.sendString((String) testPayload); + Message message1 = channel.receive(0); + assertEquals(testPayload, message1.getPayload()); + assertEquals("foo", message1.getHeaders().get("foo")); + assertEquals("bar", message1.getHeaders().get("bar")); + assertNull(message1.getHeaders().get(MessageHeaders.PREFIX + "baz")); + testPayload = 123; - gateway.sendInteger((Integer) testPayload); - Mockito.verify(handler, Mockito.times(1)).handleMessage(Mockito.any(Message.class)); - - this.prepareHandlerForTest(handler); + gatewayWithHeaderValues.sendInteger((Integer) testPayload); + Message message2 = channel.receive(0); + assertEquals(testPayload, message2.getPayload()); + assertEquals("foo", message2.getHeaders().get("foo")); + assertEquals("bar", message2.getHeaders().get("bar")); + assertNull(message2.getHeaders().get(MessageHeaders.PREFIX + "baz")); + testPayload = "withAnnotatedHeaders"; - gateway.sendStringWithParameterHeaders((String) testPayload, "headerA", "headerB"); - Mockito.verify(handler, Mockito.times(1)).handleMessage(Mockito.any(Message.class)); + gatewayWithHeaderValues.sendStringWithParameterHeaders((String) testPayload, "headerA", "headerB"); + Message message3 = channel.receive(0); + assertEquals("foo", message3.getHeaders().get("foo")); + assertEquals("bar", message3.getHeaders().get("bar")); + assertEquals("headerA", message3.getHeaders().get("headerA")); + assertEquals("headerB", message3.getHeaders().get("headerB")); } + @Test + public void validateHeaderExpressionMappings() throws Exception { + testPayload = "hello"; + gatewayWithHeaderExpressions.sendString((String) testPayload); + Message message1 = channel.receive(0); + assertEquals(testPayload, message1.getPayload()); + assertEquals(42, message1.getHeaders().get("foo")); + assertEquals("foobar", message1.getHeaders().get("bar")); + assertNull(message1.getHeaders().get(MessageHeaders.PREFIX + "baz")); - private void prepareHandlerForTest(MessageHandler handler) { - Mockito.reset(handler); - Mockito.doAnswer(new Answer() { - public Object answer(InvocationOnMock invocation) { - Message message = (Message) invocation.getArguments()[0]; - assertEquals(testPayload, message.getPayload()); - assertEquals("foo", message.getHeaders().get("foo")); - assertEquals("bar", message.getHeaders().get("bar")); - assertNull(message.getHeaders().get(MessageHeaders.PREFIX + "baz")); - if (message.getPayload().equals("withAnnotatedHeaders")){ - assertEquals("headerA", message.getHeaders().get("headerA")); - assertEquals("headerB", message.getHeaders().get("headerB")); - } - return null; - }}) - .when(handler).handleMessage(Mockito.any(Message.class)); + testPayload = 123; + gatewayWithHeaderExpressions.sendInteger((Integer) testPayload); + Message message2 = channel.receive(0); + assertEquals(testPayload, message2.getPayload()); + assertEquals(42, message2.getHeaders().get("foo")); + assertEquals("foobar", message2.getHeaders().get("bar")); + assertNull(message2.getHeaders().get(MessageHeaders.PREFIX + "baz")); + + testPayload = "withAnnotatedHeaders"; + gatewayWithHeaderExpressions.sendStringWithParameterHeaders((String) testPayload, "headerA", "headerB"); + Message message3 = channel.receive(0); + assertEquals(42, message3.getHeaders().get("foo")); + assertEquals("foobar", message3.getHeaders().get("bar")); + assertEquals("headerA", message3.getHeaders().get("headerA")); + assertEquals("headerB", message3.getHeaders().get("headerB")); } diff --git a/spring-integration-core/src/test/java/org/springframework/integration/handler/ArgumentArrayMessageMapperToMessageTests.java b/spring-integration-core/src/test/java/org/springframework/integration/handler/ArgumentArrayMessageMapperToMessageTests.java index abc7236bbb..cfd0983da7 100644 --- a/spring-integration-core/src/test/java/org/springframework/integration/handler/ArgumentArrayMessageMapperToMessageTests.java +++ b/spring-integration-core/src/test/java/org/springframework/integration/handler/ArgumentArrayMessageMapperToMessageTests.java @@ -25,6 +25,9 @@ import java.util.Map; import org.junit.Test; +import org.springframework.expression.Expression; +import org.springframework.expression.common.LiteralExpression; +import org.springframework.expression.spel.standard.SpelExpressionParser; import org.springframework.integration.Message; import org.springframework.integration.MessageHeaders; import org.springframework.integration.annotation.Header; @@ -192,17 +195,17 @@ public class ArgumentArrayMessageMapperToMessageTests { } @Test - public void toMessageWithPayloadAndStaticHeaders() throws Exception { + public void toMessageWithPayloadAndHeaders() throws Exception { Method method = TestService.class.getMethod("sendPayload", String.class); - Map headers = new HashMap(); - headers.put("foo", "foo"); - headers.put("bar", "bar"); - headers.put(MessageHeaders.PREFIX + "baz", "hello"); + Map headers = new HashMap(); + headers.put("foo", new LiteralExpression("foo")); + headers.put("bar", new SpelExpressionParser().parseExpression("6 * 7")); + headers.put(MessageHeaders.PREFIX + "baz", new LiteralExpression("hello")); ArgumentArrayMessageMapper mapper = new ArgumentArrayMessageMapper(method, headers); Message message = mapper.toMessage(new Object[] { "test" }); assertEquals("test", message.getPayload()); assertEquals("foo", message.getHeaders().get("foo")); - assertEquals("bar", message.getHeaders().get("bar")); + assertEquals(42, message.getHeaders().get("bar")); }