INT-1354 added support for "expression" attributes on the <header> sub-elements of a <gateway>

This commit is contained in:
Mark Fisher
2010-08-17 19:15:27 +00:00
parent 700b89e99a
commit 1f0d9d6e45
7 changed files with 149 additions and 78 deletions

View File

@@ -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<Element> invocationHeaders){
Map<String, Object> methodInvocationHeaders = new ManagedMap<String, Object>();
private void setMethodInvocationHeaders(BeanDefinitionBuilder gatewayDefinitionBuilder, List<Element> invocationHeaders) {
Map<String, Object> headerExpressions = new ManagedMap<String, Object>();
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);
}
}

View File

@@ -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.
* &lt;si:method name="echo" request-channel="inputA" reply-timeout="2" request-timeout="200"/&gt;
@@ -38,7 +40,7 @@ public class GatewayMethodDefinition {
private volatile String replyTimeout;
private volatile Map<String, Object> staticHeaders = new HashMap<String, Object>();
private volatile Map<String, Expression> headerExpressions = new HashMap<String, Expression>();
public String getPayloadExpression() {
@@ -49,12 +51,12 @@ public class GatewayMethodDefinition {
this.payloadExpression = payloadExpression;
}
public Map<String, Object> getStaticHeaders() {
return staticHeaders;
public Map<String, Expression> getHeaderExpressions() {
return this.headerExpressions;
}
public void setStaticHeaders(Map<String, Object> staticHeaders) {
this.staticHeaders = staticHeaders;
public void setHeaderExpressions(Map<String, Expression> headerExpressions) {
this.headerExpressions = headerExpressions;
}
public String getRequestChannelName() {

View File

@@ -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<String, Object> staticHeaders = null;
Map<String, Expression> 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);
}

View File

@@ -108,7 +108,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper<Object[]
private final Method method;
private final Map<String, Object> staticHeaders;
private final Map<String, Expression> headerExpressions;
private final List<MethodParameter> parameterList;
@@ -116,7 +116,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper<Object[]
private final Map<String, Expression> parameterPayloadExpressions = new HashMap<String, Expression>();
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<Object[]
this(method, null);
}
public ArgumentArrayMessageMapper(Method method, Map<String, Object> staticHeaders) {
public ArgumentArrayMessageMapper(Method method, Map<String, Expression> 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<Object[]
public void setBeanFactory(final BeanFactory beanFactory) {
if (beanFactory != null) {
this.beanResolver = new SimpleBeanResolver(beanFactory);
this.staticEvaluationContext.setBeanResolver(beanResolver);
this.evaluationContext.setBeanResolver(beanResolver);
}
}
@@ -155,7 +155,6 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper<Object[]
return this.mapArgumentsToMessage(arguments);
}
@SuppressWarnings("unchecked")
private Message<?> mapArgumentsToMessage(Object[] arguments) {
Object messageOrPayload = null;
boolean foundPayloadAnnotation = false;
@@ -199,10 +198,10 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper<Object[]
if (!(argumentValue instanceof Map)) {
throw new IllegalArgumentException("@Headers annotation is only valid for Map-typed parameters");
}
for (Object key : ((Map) argumentValue).keySet()) {
for (Object key : ((Map<?, ?>) 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<Object[]
throw new MessagingException("Ambiguous method parameters; found more than one " +
"Map-typed parameter and neither one contains a @Payload annotation");
}
this.copyHeaders((Map) argumentValue, headers);
this.copyHeaders((Map<?, ?>) 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<String, Object> evaluatedHeaders = new HashMap<String, Object>();
for (Map.Entry<String, Expression> 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<Object[]
expression = PARSER.parseExpression(expressionString);
this.parameterPayloadExpressions.put(expressionString, expression);
}
return expression.getValue(this.staticEvaluationContext, argumentValue);
return expression.getValue(this.evaluationContext, argumentValue);
}
private Annotation findMappingAnnotation(Annotation[] annotations) {
@@ -260,8 +266,7 @@ public class ArgumentArrayMessageMapper implements InboundMessageMapper<Object[]
return match;
}
@SuppressWarnings("unchecked")
private void copyHeaders(Map argumentValue, Map<String, Object> headers) {
private void copyHeaders(Map<?, ?> argumentValue, Map<String, Object> headers) {
for (Object key : argumentValue.keySet()) {
if (!(key instanceof String)) {
throw new IllegalArgumentException("Invalid header name [" + key +

View File

@@ -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">
<int:gateway id="gateway"
<int:gateway id="gatewayWithHeaderValues"
service-interface="org.springframework.integration.gateway.HeaderEnrichedGatewayTests$SampleGateway">
<int:method name="sendString" request-channel="input">
<int:header name="foo" value="#{stringValue}"/>
<int:method name="sendString" request-channel="channel">
<int:header name="foo" value="#{stringValueFoo}"/>
<int:header name="bar" value="bar"/>
</int:method>
<int:method name="sendInteger" request-channel="input">
<int:method name="sendInteger" request-channel="channel">
<int:header name="foo" value="foo"/>
<int:header name="bar" value="bar"/>
</int:method>
<int:method name="sendStringWithParameterHeaders" request-channel="input">
<int:method name="sendStringWithParameterHeaders" request-channel="channel">
<int:header name="foo" value="foo"/>
<int:header name="bar" value="bar"/>
</int:method>
</int:gateway>
<int:gateway id="gatewayWithHeaderExpressions"
service-interface="org.springframework.integration.gateway.HeaderEnrichedGatewayTests$SampleGateway">
<int:method name="sendString" request-channel="channel">
<int:header name="foo" expression="6 * 7"/>
<int:header name="bar" expression="'foobar'"/>
</int:method>
<int:method name="sendInteger" request-channel="channel">
<int:header name="foo" expression="42"/>
<int:header name="bar" expression="@stringValueFoo + @stringValueBar"/>
</int:method>
<int:method name="sendStringWithParameterHeaders" request-channel="channel">
<int:header name="foo" expression="@stringValueFoo.length() + 39"/>
<int:header name="bar" expression="'foo' + @stringValueBar"/>
</int:method>
</int:gateway>
<bean id="stringValue" class="java.lang.String">
<bean id="stringValueFoo" class="java.lang.String">
<constructor-arg value="foo"/>
</bean>
<bean id="stringValueBar" class="java.lang.String">
<constructor-arg value="bar"/>
</bean>
<int:channel id="input"/>
<int:channel id="channel">
<int:queue/>
</int:channel>
</beans>

View File

@@ -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<Object>() {
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"));
}

View File

@@ -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<String, Object> headers = new HashMap<String, Object>();
headers.put("foo", "foo");
headers.put("bar", "bar");
headers.put(MessageHeaders.PREFIX + "baz", "hello");
Map<String, Expression> headers = new HashMap<String, Expression>();
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"));
}