diff --git a/spring-rabbit/src/main/java/org/springframework/amqp/rabbit/support/PublisherCallbackChannelImpl.java b/spring-rabbit/src/main/java/org/springframework/amqp/rabbit/support/PublisherCallbackChannelImpl.java index d3c9a2a5..193f3727 100644 --- a/spring-rabbit/src/main/java/org/springframework/amqp/rabbit/support/PublisherCallbackChannelImpl.java +++ b/spring-rabbit/src/main/java/org/springframework/amqp/rabbit/support/PublisherCallbackChannelImpl.java @@ -17,12 +17,15 @@ package org.springframework.amqp.rabbit.support; import java.io.IOException; import java.util.Collections; +import java.util.HashSet; import java.util.Iterator; import java.util.Map; import java.util.Map.Entry; +import java.util.Set; import java.util.SortedMap; import java.util.TreeMap; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentSkipListMap; import java.util.concurrent.TimeoutException; import org.apache.commons.logging.Log; @@ -72,7 +75,7 @@ public class PublisherCallbackChannelImpl implements PublisherCallbackChannel, C private final Map> pendingConfirms = new ConcurrentHashMap>(); - private final Map listenerForSeq = new ConcurrentHashMap(); + private final SortedMap listenerForSeq = new ConcurrentSkipListMap(); public PublisherCallbackChannelImpl(Channel delegate) { this.delegate = delegate; @@ -483,27 +486,55 @@ public class PublisherCallbackChannelImpl implements PublisherCallbackChannel, C } private void processAck(long seq, boolean ack, boolean multiple) { - Listener listener = this.listenerForSeq.get(seq); - if (listener != null && listener.isConfirmListener()) { - if (multiple) { - Map headMap = this.pendingConfirms.get(listener).headMap(seq + 1); - synchronized(this.pendingConfirms) { - Iterator> iterator = headMap.entrySet().iterator(); - while (iterator.hasNext()) { - Entry entry = iterator.next(); - iterator.remove(); - listener.handleConfirm(entry.getValue(), ack); + if (multiple) { + /* + * Piggy-backed ack - extract all Listeners for this and earlier + * sequences. Then, for each Listener, handle each of it's acks. + */ + synchronized(this.pendingConfirms) { + Map involvedListeners = this.listenerForSeq.headMap(seq + 1); + // eliminate duplicates + Set listeners = new HashSet(involvedListeners.values()); + for (Listener involvedListener : listeners) { + // find all unack'd confirms for this listener and handle them + SortedMap confirmsMap = this.pendingConfirms.get(involvedListener); + if (confirmsMap != null) { + Map confirms = confirmsMap.headMap(seq + 1); + Iterator> iterator = confirms.entrySet().iterator(); + while (iterator.hasNext()) { + Entry entry = iterator.next(); + iterator.remove(); + doHandleConfirm(ack, involvedListener, entry.getValue()); + } } } } - else { + } + else { + Listener listener = this.listenerForSeq.get(seq); + if (listener != null) { PendingConfirm pendingConfirm = this.pendingConfirms.get(listener).remove(seq); if (pendingConfirm != null) { - listener.handleConfirm(pendingConfirm, ack); + doHandleConfirm(ack, listener, pendingConfirm); } } - } else { - logger.error("No listener for seq:" + seq); + else { + logger.error("No listener for seq:" + seq); + } + } + } + + private void doHandleConfirm(boolean ack, Listener listener, PendingConfirm pendingConfirm) { + try { + if (listener.isConfirmListener()) { + if (logger.isDebugEnabled()) { + logger.debug("Sending confirm " + pendingConfirm); + } + listener.handleConfirm(pendingConfirm, ack); + } + } + catch (Exception e) { + logger.error("Exception delivering confirm", e); } } diff --git a/spring-rabbit/src/test/java/org/springframework/amqp/rabbit/core/RabbitTemplatePublisherCallbacksIntegrationTests.java b/spring-rabbit/src/test/java/org/springframework/amqp/rabbit/core/RabbitTemplatePublisherCallbacksIntegrationTests.java index 6fb7c920..a951e5d4 100644 --- a/spring-rabbit/src/test/java/org/springframework/amqp/rabbit/core/RabbitTemplatePublisherCallbacksIntegrationTests.java +++ b/spring-rabbit/src/test/java/org/springframework/amqp/rabbit/core/RabbitTemplatePublisherCallbacksIntegrationTests.java @@ -366,4 +366,67 @@ public class RabbitTemplatePublisherCallbacksIntegrationTests { Collection unconfirmed = template.getUnconfirmed(0); assertNull(unconfirmed); } + + /** + * Tests that piggy-backed confirms (multiple=true) are distributed to the proper + * template. + * @throws Exception + */ + @Test + public void testPublisherConfirmMultipleWithTwoListeners() throws Exception { + ConnectionFactory mockConnectionFactory = mock(ConnectionFactory.class); + Connection mockConnection = mock(Connection.class); + Channel mockChannel = mock(Channel.class); + + when(mockConnectionFactory.newConnection((ExecutorService) null)).thenReturn(mockConnection); + when(mockConnection.isOpen()).thenReturn(true); + PublisherCallbackChannelImpl callbackChannel = new PublisherCallbackChannelImpl(mockChannel); + when(mockConnection.createChannel()).thenReturn(callbackChannel); + + final AtomicInteger count = new AtomicInteger(); + doAnswer(new Answer(){ + public Object answer(InvocationOnMock invocation) throws Throwable { + return count.incrementAndGet(); + }}).when(mockChannel).getNextPublishSeqNo(); + + final RabbitTemplate template1 = new RabbitTemplate(new SingleConnectionFactory(mockConnectionFactory)); + + final Set confirms = new HashSet(); + final CountDownLatch latch1 = new CountDownLatch(1); + template1.setConfirmCallback(new ConfirmCallback() { + + public void confirm(CorrelationData correlationData, boolean ack) { + if (ack) { + confirms.add(correlationData.getId() + "1"); + latch1.countDown(); + } + } + }); + final RabbitTemplate template2 = new RabbitTemplate(new SingleConnectionFactory(mockConnectionFactory)); + + final CountDownLatch latch2 = new CountDownLatch(1); + template2.setConfirmCallback(new ConfirmCallback() { + + public void confirm(CorrelationData correlationData, boolean ack) { + if (ack) { + confirms.add(correlationData.getId() + "2"); + latch2.countDown(); + } + } + }); + template1.convertAndSend(ROUTE, (Object) "message", new CorrelationData("abc")); + template2.convertAndSend(ROUTE, (Object) "message", new CorrelationData("def")); + template2.convertAndSend(ROUTE, (Object) "message", new CorrelationData("ghi")); + callbackChannel.handleAck(3, true); + assertTrue(latch1.await(1000, TimeUnit.MILLISECONDS)); + assertTrue(latch2.await(1000, TimeUnit.MILLISECONDS)); + Collection unconfirmed1 = template1.getUnconfirmed(0); + assertNull(unconfirmed1); + Collection unconfirmed2 = template2.getUnconfirmed(0); + assertNull(unconfirmed2); + assertTrue(confirms.contains("abc1")); + assertTrue(confirms.contains("def2")); + assertTrue(confirms.contains("ghi2")); + assertEquals(3, confirms.size()); + } }