Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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);
}

}

Expand Down Expand Up @@ -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) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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) {
Expand All @@ -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<MessageQueue> mqs, Message msg, Object arg) {
return null;
}
};

Expand All @@ -506,14 +499,10 @@ public MessageQueue select(List<MessageQueue> mqs, Message msg, Object arg) {
failBecauseExceptionWasNotThrown(Exception.class);
} catch (Exception e) {
ConcurrentHashMap<String, RequestResponseFuture> responseMap = RequestFutureHolder.getInstance().getRequestFutureTable();
assertThat(responseMap).isNotNull();
for (Map.Entry<String, RequestResponseFuture> 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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -75,6 +77,8 @@
@RunWith(MockitoJUnitRunner.class)
public class DefaultMQProducerImplTest {

private static final String CORRELATION_ID = "correlation-id";

@Mock
private Message message;

Expand Down Expand Up @@ -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));
Expand All @@ -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));
Expand All @@ -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
Expand Down
Loading