MessageBus now manages the subscription of a MessageEndpoint to a SynchronousChannel if that endpoint's Subscription object contains a NULL Schedule.

This commit is contained in:
Mark Fisher
2008-04-04 20:56:12 +00:00
parent 5821ac23a2
commit 8ea5f012f3
7 changed files with 203 additions and 37 deletions

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2007 the original author or authors.
* Copyright 2002-2008 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.
@@ -41,6 +41,7 @@ import org.springframework.integration.channel.MessageChannel;
import org.springframework.integration.channel.SimpleChannel;
import org.springframework.integration.dispatcher.DefaultMessageDispatcher;
import org.springframework.integration.dispatcher.SchedulingMessageDispatcher;
import org.springframework.integration.dispatcher.SynchronousChannel;
import org.springframework.integration.endpoint.ConcurrencyPolicy;
import org.springframework.integration.endpoint.DefaultEndpointRegistry;
import org.springframework.integration.endpoint.DefaultMessageEndpoint;
@@ -125,7 +126,6 @@ public class MessageBus implements ChannelRegistry, EndpointRegistry, Applicatio
* is registered without an explicitly provided policy of its own.
*/
public void setDefaultConcurrencyPolicy(ConcurrencyPolicy defaultConcurrencyPolicy) {
Assert.notNull(defaultConcurrencyPolicy, "'defaultConcurrencyPolicy' must not be null");
this.defaultConcurrencyPolicy = defaultConcurrencyPolicy;
}
@@ -179,9 +179,6 @@ public class MessageBus implements ChannelRegistry, EndpointRegistry, Applicatio
return;
}
this.initializing = true;
if (this.defaultConcurrencyPolicy == null) {
this.defaultConcurrencyPolicy = new ConcurrencyPolicy();
}
if (this.executor == null) {
this.executor = new ScheduledThreadPoolExecutor(DEFAULT_DISPATCHER_POOL_SIZE);
}
@@ -257,7 +254,8 @@ public class MessageBus implements ChannelRegistry, EndpointRegistry, Applicatio
if (endpoint instanceof ChannelRegistryAware) {
((ChannelRegistryAware) endpoint).setChannelRegistry(this.channelRegistry);
}
if (endpoint.getConcurrencyPolicy() == null && endpoint instanceof DefaultMessageEndpoint) {
if (endpoint.getConcurrencyPolicy() == null && this.defaultConcurrencyPolicy != null
&& endpoint instanceof DefaultMessageEndpoint) {
((DefaultMessageEndpoint) endpoint).setConcurrencyPolicy(this.defaultConcurrencyPolicy);
}
this.endpointRegistry.registerEndpoint(name, endpoint);
@@ -366,6 +364,13 @@ public class MessageBus implements ChannelRegistry, EndpointRegistry, Applicatio
}
private void registerWithDispatcher(MessageChannel channel, MessageHandler handler, Schedule schedule) {
if (schedule == null && (channel instanceof SynchronousChannel)) {
((SynchronousChannel) channel).addHandler(handler);
if (handler instanceof Lifecycle) {
((Lifecycle) handler).start();
}
return;
}
SchedulingMessageDispatcher dispatcher = dispatchers.get(channel);
if (dispatcher == null) {
if (logger.isWarnEnabled()) {

View File

@@ -57,6 +57,10 @@ public class DefaultMessageDistributor implements MessageDistributor {
this.handlers.add(handler);
}
public boolean removeHandler(MessageHandler handler) {
return this.handlers.remove(handler);
}
public boolean distribute(Message<?> message) {
int attempts = 0;
List<MessageHandler> targets = new ArrayList<MessageHandler>(this.handlers);

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2007 the original author or authors.
* Copyright 2002-2008 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.
@@ -29,6 +29,8 @@ public interface MessageDistributor {
void addHandler(MessageHandler handler);
boolean removeHandler(MessageHandler handler);
boolean distribute(Message<?> message);
}

View File

@@ -19,6 +19,9 @@ package org.springframework.integration.dispatcher;
import java.util.ArrayList;
import java.util.Collection;
import java.util.List;
import java.util.Queue;
import java.util.concurrent.LinkedBlockingQueue;
import java.util.concurrent.atomic.AtomicInteger;
import org.springframework.integration.adapter.PollableSource;
import org.springframework.integration.channel.AbstractMessageChannel;
@@ -30,38 +33,56 @@ import org.springframework.integration.message.selector.MessageSelector;
import org.springframework.util.CollectionUtils;
/**
* A channel that invokes any subscribed {@link MessageHandler handler(s)} in a
* sender's thread. If a {@link PollableSource} is provided, then that source
* will likewise be polled within a receiver's thread. If no source is provided,
* then receive() will always return null.
* A channel that invokes the subscribed {@link MessageHandler handler(s)} in a
* sender's thread (returning after at most one handles the message). If a
* {@link PollableSource} is provided, then that source will likewise be polled
* within a receiver's thread.
* <p>
* If the channel has no subscribed handlers and no configured source, then it
* will store messages in a thread-bound queue. In other words, send() will put
* a message at the tail of the queue for the current thread, and receive() will
* retrieve a message from the head of the queue.
*
* @author Dave Syer
* @author Mark Fisher
*/
public class SynchronousChannel extends AbstractMessageChannel {
private static final ThreadLocalMessageHolder messageHolder = new ThreadLocalMessageHolder();
private final PollableSource<?> source;
private final MessageDistributor distributor;
private volatile PollableSource<?> source;
private final AtomicInteger handlerCount = new AtomicInteger();
public SynchronousChannel() {
this(null);
}
public SynchronousChannel(DispatcherPolicy dispatcherPolicy) {
super(dispatcherPolicy != null ? dispatcherPolicy : new DispatcherPolicy());
public SynchronousChannel(PollableSource<?> source) {
super(defaultDispatcherPolicy());
this.source = source;
this.distributor = new DefaultMessageDistributor(this.getDispatcherPolicy());
}
public void setSource(PollableSource<?> source) {
this.source = source;
}
public void addHandler(MessageHandler handler) {
this.distributor.addHandler(handler);
this.handlerCount.incrementAndGet();
}
public boolean removeHandler(MessageHandler handler) {
if (this.distributor.removeHandler(handler)) {
this.handlerCount.decrementAndGet();
return true;
}
return false;
}
@Override
protected Message<?> doReceive(long timeout) {
if (this.source != null) {
@@ -72,20 +93,70 @@ public class SynchronousChannel extends AbstractMessageChannel {
new SimplePayloadMessageMapper<Object>().toMessage(result);
}
}
return null;
return messageHolder.get().poll();
}
@Override
protected boolean doSend(Message<?> message, long timeout) {
return this.distributor.distribute(message);
if (message == null) {
return false;
}
if (this.handlerCount.get() > 0) {
return this.distributor.distribute(message);
}
else if (this.source == null) {
return messageHolder.get().add(message);
}
return false;
}
/**
* Remove and return any messages that are stored for the current thread.
*/
public List<Message<?>> clear() {
return new ArrayList<Message<?>>();
List<Message<?>> removedMessages = new ArrayList<Message<?>>();
Message<?> next = messageHolder.get().poll();
while (next != null) {
removedMessages.add(next);
next = messageHolder.get().poll();
}
return removedMessages;
}
/**
* Remove and return any messages that are stored for the current thread
* and do not match the provided selector.
*/
public List<Message<?>> purge(MessageSelector selector) {
return new ArrayList<Message<?>>();
List<Message<?>> removedMessages = new ArrayList<Message<?>>();
Object[] allMessages = messageHolder.get().toArray();
for (Object next : allMessages) {
Message<?> message = (Message<?>) next;
if (!selector.accept(message) && messageHolder.get().remove(message)) {
removedMessages.add(message);
}
}
return removedMessages;
}
private static DispatcherPolicy defaultDispatcherPolicy() {
DispatcherPolicy dispatcherPolicy = new DispatcherPolicy(false);
dispatcherPolicy.setMaxMessagesPerTask(1);
dispatcherPolicy.setReceiveTimeout(0);
dispatcherPolicy.setRejectionLimit(1);
dispatcherPolicy.setRetryInterval(0);
dispatcherPolicy.setShouldFailOnRejectionLimit(false);
return dispatcherPolicy;
}
private static class ThreadLocalMessageHolder extends ThreadLocal<Queue<Message<?>>> {
@Override
protected Queue<Message<?>> initialValue() {
return new LinkedBlockingQueue<Message<?>>();
}
}
}

View File

@@ -0,0 +1,71 @@
/*
* Copyright 2002-2008 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.bus;
import static org.junit.Assert.assertEquals;
import org.junit.Before;
import org.junit.Test;
import org.springframework.integration.channel.MessageChannel;
import org.springframework.integration.dispatcher.SynchronousChannel;
import org.springframework.integration.endpoint.DefaultMessageEndpoint;
import org.springframework.integration.handler.MessageHandler;
import org.springframework.integration.message.Message;
import org.springframework.integration.message.StringMessage;
import org.springframework.integration.scheduling.Subscription;
/**
* @author Mark Fisher
*/
public class SynchronousChannelSubscriptionTests {
private MessageBus bus = new MessageBus();
private MessageChannel sourceChannel = new SynchronousChannel();
private MessageChannel targetChannel = new SynchronousChannel();
@Before
public void setupChannels() {
bus.registerChannel("sourceChannel", sourceChannel);
bus.registerChannel("targetChannel", targetChannel);
}
@Test
public void testSendAndReceive() throws InterruptedException {
DefaultMessageEndpoint endpoint = new DefaultMessageEndpoint(new TestHandler());
endpoint.setSubscription(new Subscription("sourceChannel"));
endpoint.setDefaultOutputChannelName("targetChannel");
bus.registerEndpoint("testEndpoint", endpoint);
bus.start();
this.sourceChannel.send(new StringMessage("foo"));
Message<?> response = this.targetChannel.receive();
assertEquals("foo!", response.getPayload());
}
private static class TestHandler implements MessageHandler {
public Message<?> handle(Message<?> message) {
return new StringMessage(message.getPayload() + "!");
}
}
}

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2007 the original author or authors.
* Copyright 2002-2008 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.
@@ -25,7 +25,6 @@ import org.springframework.beans.DirectFieldAccessor;
import org.springframework.context.ApplicationContext;
import org.springframework.context.support.ClassPathXmlApplicationContext;
import org.springframework.integration.bus.MessageBus;
import org.springframework.integration.endpoint.ConcurrentHandler;
import org.springframework.integration.endpoint.DefaultMessageEndpoint;
import org.springframework.integration.handler.MessageHandlerChain;
import org.springframework.integration.router.AggregatingMessageHandler;
@@ -85,10 +84,7 @@ public class AggregatorAnnotationTests {
final String endpointName) {
MessageBus messageBus = getMessageBus(context);
DefaultMessageEndpoint endpoint = (DefaultMessageEndpoint) messageBus.lookupEndpoint(endpointName + "-endpoint");
ConcurrentHandler handler = (ConcurrentHandler) endpoint.getHandler();
DirectFieldAccessor concurrentHandlerAccessor = new DirectFieldAccessor(handler);
MessageHandlerChain messageHandlerChain = (MessageHandlerChain) concurrentHandlerAccessor
.getPropertyValue("handler");
MessageHandlerChain messageHandlerChain = (MessageHandlerChain) endpoint.getHandler();
AggregatingMessageHandler aggregatingMessageHandler = (AggregatingMessageHandler) ((List) new DirectFieldAccessor(
messageHandlerChain).getPropertyValue("handlers")).get(0);
DirectFieldAccessor aggregatingMessageHandlerAccessor = new DirectFieldAccessor(aggregatingMessageHandler);

View File

@@ -17,13 +17,13 @@
package org.springframework.integration.dispatcher;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertFalse;
import static org.junit.Assert.assertNotNull;
import static org.junit.Assert.assertNull;
import static org.junit.Assert.assertTrue;
import java.util.Collection;
import java.util.Collections;
import java.util.List;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.SynchronousQueue;
import java.util.concurrent.TimeUnit;
@@ -54,10 +54,30 @@ public class SynchronousChannelTests {
}
@Test
public void testSendWithNoHandler() {
public void testSendAndReceiveWithNoHandler() {
SynchronousChannel channel = new SynchronousChannel();
StringMessage message = new StringMessage("test");
assertFalse(channel.send(message));
assertNull(channel.receive());
assertTrue(channel.send(message));
Message<?> response = channel.receive();
assertNotNull(response);
assertEquals(response, message);
assertNull(channel.receive());
}
@Test
public void testSendAndClearWithNoHandler() {
SynchronousChannel channel = new SynchronousChannel();
StringMessage message1 = new StringMessage("test1");
StringMessage message2 = new StringMessage("test2");
assertNull(channel.receive());
assertTrue(channel.send(message1));
assertTrue(channel.send(message2));
List<Message<?>> clearedMessages = channel.clear();
assertEquals(2, clearedMessages.size());
assertEquals(message1, clearedMessages.get(0));
assertEquals(message2, clearedMessages.get(1));
assertNull(channel.receive());
}
@Test
@@ -78,8 +98,7 @@ public class SynchronousChannelTests {
@Test
public void testReceive() {
SynchronousChannel channel = new SynchronousChannel();
channel.setSource(new PollableSource<String>() {
SynchronousChannel channel = new SynchronousChannel(new PollableSource<String>() {
public Collection<String> poll(int limit) {
return Collections.singleton("foo");
}
@@ -93,8 +112,7 @@ public class SynchronousChannelTests {
@Test
public void testReceiveWithMessageResult() {
SynchronousChannel channel = new SynchronousChannel();
channel.setSource(new MessageReturningTestSource("foo"));
SynchronousChannel channel = new SynchronousChannel(new MessageReturningTestSource("foo"));
Message<?> message = channel.receive();
assertNotNull(message);
assertNotNull(message.getPayload());
@@ -106,8 +124,7 @@ public class SynchronousChannelTests {
@Test
public void testReceiveInSeparateThread() throws InterruptedException {
final SynchronousChannel channel = new SynchronousChannel();
channel.setSource(new MessageReturningTestSource("foo"));
final SynchronousChannel channel = new SynchronousChannel(new MessageReturningTestSource("foo"));
final SynchronousQueue<Message<?>> messageHolder = new SynchronousQueue<Message<?>>();
new Thread(new Runnable() {
public void run() {