diff --git a/client/src/main/java/org/apache/rocketmq/client/impl/producer/DefaultMQProducerImpl.java b/client/src/main/java/org/apache/rocketmq/client/impl/producer/DefaultMQProducerImpl.java index 9ad5fcef4dc..a8c53b2fc5a 100644 --- a/client/src/main/java/org/apache/rocketmq/client/impl/producer/DefaultMQProducerImpl.java +++ b/client/src/main/java/org/apache/rocketmq/client/impl/producer/DefaultMQProducerImpl.java @@ -1659,22 +1659,31 @@ public void request(Message msg, final RequestCallback requestCallback, long tim RequestFutureHolder.getInstance().getRequestFutureTable().put(correlationId, requestResponseFuture); long cost = System.currentTimeMillis() - beginTimestamp; - this.sendDefaultImpl(msg, CommunicationMode.ASYNC, new SendCallback() { - @Override - public void onSuccess(SendResult sendResult) { - // Only mark the request as sent here. The user callback must fire when the reply - // arrives (processReplyMessage), on timeout (scanExpiredRequest), or on send - // failure (requestFail). Invoking it here delivers a premature onSuccess(null) - // and, combined with the later reply/timeout callback, causes a double callback. - requestResponseFuture.setSendRequestOk(true); - } + boolean sendInvocationCompleted = false; + try { + this.sendDefaultImpl(msg, CommunicationMode.ASYNC, new SendCallback() { + @Override + public void onSuccess(SendResult sendResult) { + // Only mark the request as sent here. The user callback must fire when the reply + // arrives (processReplyMessage), on timeout (scanExpiredRequest), or on send + // failure (requestFail). Invoking it here delivers a premature onSuccess(null) + // and, combined with the later reply/timeout callback, causes a double callback. + requestResponseFuture.setSendRequestOk(true); + } - @Override - public void onException(Throwable e) { - requestResponseFuture.setCause(e); - requestFail(correlationId); + @Override + public void onException(Throwable e) { + requestResponseFuture.setCause(e); + requestFail(correlationId); + } + }, timeout - cost); + sendInvocationCompleted = true; + } finally { + if (!sendInvocationCompleted) { + RequestFutureHolder.getInstance().getRequestFutureTable() + .remove(correlationId, requestResponseFuture); } - }, timeout - cost); + } } public Message request(final Message msg, final MessageQueueSelector selector, final Object arg, @@ -1720,18 +1729,27 @@ public void request(final Message msg, final MessageQueueSelector selector, fina RequestFutureHolder.getInstance().getRequestFutureTable().put(correlationId, requestResponseFuture); long cost = System.currentTimeMillis() - beginTimestamp; - this.sendSelectImpl(msg, selector, arg, CommunicationMode.ASYNC, new SendCallback() { - @Override - public void onSuccess(SendResult sendResult) { - requestResponseFuture.setSendRequestOk(true); - } + boolean sendInvocationCompleted = false; + try { + this.sendSelectImpl(msg, selector, arg, CommunicationMode.ASYNC, new SendCallback() { + @Override + public void onSuccess(SendResult sendResult) { + requestResponseFuture.setSendRequestOk(true); + } - @Override - public void onException(Throwable e) { - requestResponseFuture.setCause(e); - requestFail(correlationId); + @Override + public void onException(Throwable e) { + requestResponseFuture.setCause(e); + requestFail(correlationId); + } + }, timeout - cost); + sendInvocationCompleted = true; + } finally { + if (!sendInvocationCompleted) { + RequestFutureHolder.getInstance().getRequestFutureTable() + .remove(correlationId, requestResponseFuture); } - }, timeout - cost); + } } @@ -1790,18 +1808,27 @@ public void request(final Message msg, final MessageQueue mq, final RequestCallb RequestFutureHolder.getInstance().getRequestFutureTable().put(correlationId, requestResponseFuture); long cost = System.currentTimeMillis() - beginTimestamp; - this.sendKernelImpl(msg, mq, CommunicationMode.ASYNC, new SendCallback() { - @Override - public void onSuccess(SendResult sendResult) { - requestResponseFuture.setSendRequestOk(true); - } + boolean sendInvocationCompleted = false; + try { + this.sendKernelImpl(msg, mq, CommunicationMode.ASYNC, new SendCallback() { + @Override + public void onSuccess(SendResult sendResult) { + requestResponseFuture.setSendRequestOk(true); + } - @Override - public void onException(Throwable e) { - requestResponseFuture.setCause(e); - requestFail(correlationId); + @Override + public void onException(Throwable e) { + requestResponseFuture.setCause(e); + requestFail(correlationId); + } + }, null, timeout - cost); + sendInvocationCompleted = true; + } finally { + if (!sendInvocationCompleted) { + RequestFutureHolder.getInstance().getRequestFutureTable() + .remove(correlationId, requestResponseFuture); } - }, null, timeout - cost); + } } private void requestFail(final String correlationId) { diff --git a/client/src/test/java/org/apache/rocketmq/client/producer/DefaultMQProducerTest.java b/client/src/test/java/org/apache/rocketmq/client/producer/DefaultMQProducerTest.java index 33cf0df390d..86ad3e0c324 100644 --- a/client/src/test/java/org/apache/rocketmq/client/producer/DefaultMQProducerTest.java +++ b/client/src/test/java/org/apache/rocketmq/client/producer/DefaultMQProducerTest.java @@ -32,6 +32,7 @@ import org.apache.rocketmq.common.UtilAll; import org.apache.rocketmq.common.compression.CompressionType; import org.apache.rocketmq.common.message.Message; +import org.apache.rocketmq.common.message.MessageConst; import org.apache.rocketmq.common.message.MessageExt; import org.apache.rocketmq.common.message.MessageQueue; import org.apache.rocketmq.remoting.RPCHook; @@ -479,9 +480,8 @@ public void onException(Throwable e) { } @Test - public void testAsyncRequest_OnException() throws Exception { + public void testAsyncRequest_SynchronousExceptionRemovesFuture() throws Exception { final AtomicInteger cc = new AtomicInteger(0); - final CountDownLatch countDownLatch = new CountDownLatch(1); RequestCallback requestCallback = new RequestCallback() { @Override public void onSuccess(Message message) { @@ -491,13 +491,6 @@ public void onSuccess(Message message) { @Override public void onException(Throwable e) { cc.incrementAndGet(); - countDownLatch.countDown(); - } - }; - MessageQueueSelector messageQueueSelector = new MessageQueueSelector() { - @Override - public MessageQueue select(List mqs, Message msg, Object arg) { - return null; } }; @@ -506,14 +499,10 @@ public MessageQueue select(List mqs, Message msg, Object arg) { failBecauseExceptionWasNotThrown(Exception.class); } catch (Exception e) { ConcurrentHashMap responseMap = RequestFutureHolder.getInstance().getRequestFutureTable(); - assertThat(responseMap).isNotNull(); - for (Map.Entry entry : responseMap.entrySet()) { - RequestResponseFuture future = entry.getValue(); - future.getRequestCallback().onException(e); - } + String correlationId = message.getProperty(MessageConst.PROPERTY_CORRELATION_ID); + assertThat(responseMap).doesNotContainKey(correlationId); } - countDownLatch.await(defaultTimeout, TimeUnit.MILLISECONDS); - assertThat(cc.get()).isEqualTo(1); + assertThat(cc.get()).isZero(); } @Test diff --git a/client/src/test/java/org/apache/rocketmq/client/producer/selector/DefaultMQProducerImplTest.java b/client/src/test/java/org/apache/rocketmq/client/producer/selector/DefaultMQProducerImplTest.java index 77a83af19c0..1d43ec1c077 100644 --- a/client/src/test/java/org/apache/rocketmq/client/producer/selector/DefaultMQProducerImplTest.java +++ b/client/src/test/java/org/apache/rocketmq/client/producer/selector/DefaultMQProducerImplTest.java @@ -29,6 +29,7 @@ import org.apache.rocketmq.client.impl.producer.TopicPublishInfo; import org.apache.rocketmq.client.producer.MessageQueueSelector; import org.apache.rocketmq.client.producer.RequestCallback; +import org.apache.rocketmq.client.producer.RequestFutureHolder; import org.apache.rocketmq.client.producer.SendCallback; import org.apache.rocketmq.client.producer.SendResult; import org.apache.rocketmq.client.producer.TransactionListener; @@ -42,6 +43,7 @@ import org.apache.rocketmq.common.producer.RecallMessageHandle; import org.apache.rocketmq.remoting.exception.RemotingException; import org.apache.rocketmq.remoting.protocol.header.CheckTransactionStateRequestHeader; +import org.junit.After; import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; @@ -75,6 +77,8 @@ @RunWith(MockitoJUnitRunner.class) public class DefaultMQProducerImplTest { + private static final String CORRELATION_ID = "correlation-id"; + @Mock private Message message; @@ -105,6 +109,7 @@ public class DefaultMQProducerImplTest { @Before public void init() throws Exception { + RequestFutureHolder.getInstance().getRequestFutureTable().remove(CORRELATION_ID); when(mQClientFactory.getTopicRouteTable()).thenReturn(mock(ConcurrentMap.class)); when(mQClientFactory.getClientId()).thenReturn("client-id"); when(mQClientFactory.getMQAdminImpl()).thenReturn(mock(MQAdminImpl.class)); @@ -116,7 +121,7 @@ public void init() throws Exception { when(mQClientFactory.getMQClientAPIImpl()).thenReturn(mQClientAPIImpl); when(mQClientFactory.findBrokerAddressInPublish(or(isNull(), anyString()))).thenReturn(defaultBrokerAddr); when(message.getTopic()).thenReturn(defaultTopic); - when(message.getProperty(MessageConst.PROPERTY_CORRELATION_ID)).thenReturn("correlation-id"); + when(message.getProperty(MessageConst.PROPERTY_CORRELATION_ID)).thenReturn(CORRELATION_ID); when(message.getBody()).thenReturn(new byte[1]); TransactionMQProducer producer = new TransactionMQProducer("test-producer-group"); producer.setTransactionListener(mock(TransactionListener.class)); @@ -129,15 +134,43 @@ public void init() throws Exception { defaultMQProducerImpl.setServiceState(ServiceState.RUNNING); } + @After + public void removeRequestFuture() { + RequestFutureHolder.getInstance().getRequestFutureTable().remove(CORRELATION_ID); + } + @Test public void testRequest() throws Exception { defaultMQProducerImpl.request(message, messageQueue, requestCallback, defaultTimeout); + assertTrue(RequestFutureHolder.getInstance().getRequestFutureTable().containsKey(CORRELATION_ID)); defaultMQProducerImpl.request(message, queueSelector, 1, requestCallback, defaultTimeout); + assertTrue(RequestFutureHolder.getInstance().getRequestFutureTable().containsKey(CORRELATION_ID)); } - @Test(expected = MQClientException.class) - public void testRequestMQClientExceptionByVoid() throws Exception { - defaultMQProducerImpl.request(message, requestCallback, defaultTimeout); + @Test + public void testRequestCallbackRemovesFutureWhenDefaultSendThrows() { + assertThrows(MQClientException.class, + () -> defaultMQProducerImpl.request(message, requestCallback, defaultTimeout)); + assertFalse(RequestFutureHolder.getInstance().getRequestFutureTable().containsKey(CORRELATION_ID)); + } + + @Test + public void testRequestCallbackRemovesFutureWhenSelectorSendThrows() { + when(queueSelector.select(any(), any(), any())).thenThrow(new RuntimeException("select failed")); + + assertThrows(MQClientException.class, + () -> defaultMQProducerImpl.request(message, queueSelector, 1, requestCallback, defaultTimeout)); + assertFalse(RequestFutureHolder.getInstance().getRequestFutureTable().containsKey(CORRELATION_ID)); + } + + @Test + public void testRequestCallbackRemovesFutureWhenQueueSendThrows() { + when(mQClientFactory.getBrokerNameFromMessageQueue(messageQueue)).thenReturn(defaultBrokerName); + when(mQClientFactory.findBrokerAddressInPublish(defaultBrokerName)).thenReturn(null); + + assertThrows(MQClientException.class, + () -> defaultMQProducerImpl.request(message, messageQueue, requestCallback, defaultTimeout)); + assertFalse(RequestFutureHolder.getInstance().getRequestFutureTable().containsKey(CORRELATION_ID)); } @Test