diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/adapter/ResponseWriter.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/adapter/ResponseWriter.java index 25ad9feee6..9ce43db169 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/adapter/ResponseWriter.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/adapter/ResponseWriter.java @@ -40,38 +40,34 @@ public class ResponseWriter { } public static void writeResponse(StreamObserver observer, final T response) { + if (null == response) { + return; + } if (observer instanceof ServerCallStreamObserver) { - if (response == null) { - return; - } - final ServerCallStreamObserver serverCallStreamObserver = (ServerCallStreamObserver) observer; if (serverCallStreamObserver.isCancelled()) { log.warn("client has cancelled the request. response to write: {}", response); return; } - - log.debug("start to write response. response: {}", response); - serverCallStreamObserver.onNext(response); } + log.debug("start to write response. response: {}", response); + observer.onNext(response); } public static void writeException(StreamObserver observer, final Throwable e) { + if (null == e) { + return; + } if (observer instanceof ServerCallStreamObserver) { - if (null == e) { - return; - } - final ServerCallStreamObserver serverCallStreamObserver = (ServerCallStreamObserver) observer; if (serverCallStreamObserver.isCancelled()) { log.warn("Client has cancelled the request. Exception to write", e); return; } - - log.debug("Start to write error response", e); - serverCallStreamObserver.onError(e); - serverCallStreamObserver.onCompleted(); } + log.debug("Start to write error response", e); + observer.onError(e); + observer.onCompleted(); } } diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerService.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerService.java index 56a917ad68..62b7d73597 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerService.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerService.java @@ -68,7 +68,7 @@ public class ConsumerService extends BaseService { private volatile ReadQueueSelector readQueueSelector; private volatile ReceiveMessageResultFilter receiveMessageResultFilter; - private volatile ResponseHook> receiveMessageHook; + private volatile ResponseHook receiveMessageHook; private volatile ResponseHook ackMessageHook; private volatile ResponseHook nackMessageHook; private volatile ResponseHook changeInvisibleDurationHook; @@ -90,20 +90,11 @@ public class ConsumerService extends BaseService { public void receiveMessage(Context ctx, ReceiveMessageRequest request, StreamObserver responseObserver) { - this.receiveMessage(ctx, request) - .thenAccept(responses -> ResponseWriter.write(responseObserver, responses.iterator())) - .exceptionally(e -> { - ResponseWriter.write( - responseObserver, - ReceiveMessageResponse.newBuilder().setStatus(ResponseBuilder.buildStatus(e)).build() - ); - return null; - }); - } - - protected CompletableFuture> receiveMessage(Context ctx, ReceiveMessageRequest request) { - CompletableFuture> future = new CompletableFuture<>(); - + ReceiveMessageResponseStreamObserver streamObserver = new ReceiveMessageResponseStreamObserver( + ctx, + request, + receiveMessageHook, + responseObserver); try { PopMessageRequestHeader requestHeader = this.buildPopMessageRequestHeader(ctx, request); SelectableMessageQueue messageQueue = this.readQueueSelector.select(ctx, request, requestHeader); @@ -112,22 +103,17 @@ public class ConsumerService extends BaseService { throw new ProxyException(Code.FORBIDDEN, "no readable topic route for topic " + requestHeader.getTopic()); } - future = this.readConsumer.popMessage( + this.readConsumer.popMessage( ctx, messageQueue.getBrokerAddr(), messageQueue.getBrokerName(), requestHeader, requestHeader.getPollTime()) - .thenApply(result -> convertToReceiveMessageResponse(ctx, request, result)); + .thenAccept(result -> writeReceiveMessageResponse(ctx, request, result, streamObserver)) + .exceptionally(e -> writeReceiveMessageResponse(ctx, e, streamObserver)); } catch (Throwable t) { - future.completeExceptionally(t); + writeReceiveMessageResponse(ctx, t, streamObserver); } - future.whenComplete((response, throwable) -> { - if (receiveMessageHook != null) { - receiveMessageHook.beforeResponse(ctx, request, response, throwable); - } - }); - return future; } protected PopMessageRequestHeader buildPopMessageRequestHeader(Context ctx, ReceiveMessageRequest request) { @@ -136,40 +122,52 @@ public class ConsumerService extends BaseService { return GrpcConverter.buildPopMessageRequestHeader(request, GrpcConverter.buildPollTimeFromContext(ctx), fifo); } - protected List convertToReceiveMessageResponse(Context ctx, ReceiveMessageRequest request, - PopResult result) { - List responseList = new ArrayList<>(); + protected void writeReceiveMessageResponse(Context ctx, ReceiveMessageRequest request, + PopResult result, StreamObserver streamObserver) { PopStatus status = result.getPopStatus(); - switch (status) { - case FOUND: - List messageList = this.receiveMessageResultFilter.filterMessage(ctx, request, result.getMsgFoundList()); - if (messageList.isEmpty()) { - responseList.add(ReceiveMessageResponse.newBuilder() + try { + switch (status) { + case FOUND: + List messageList = this.receiveMessageResultFilter.filterMessage(ctx, request, result.getMsgFoundList()); + if (messageList.isEmpty()) { + streamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) + .build()); + } else { + streamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.OK, Code.OK.name())) + .build()); + for (Message message : messageList) { + streamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setMessage(message) + .build()); + } + } + break; + case POLLING_FULL: + streamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.TOO_MANY_REQUESTS, "polling full")) + .build()); + break; + case NO_NEW_MSG: + case POLLING_NOT_FOUND: + default: + streamObserver.onNext(ReceiveMessageResponse.newBuilder() .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) .build()); - } else { - for (Message message : messageList) { - responseList.add(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, Code.OK.name())) - .setMessage(message) - .build()); - } - } - break; - case POLLING_FULL: - responseList.add(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.TOO_MANY_REQUESTS, "polling full")) - .build()); - break; - case NO_NEW_MSG: - case POLLING_NOT_FOUND: - default: - responseList.add(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) - .build()); - break; + break; + } + } finally { + streamObserver.onCompleted(); } - return responseList; + } + + protected Void writeReceiveMessageResponse(Context ctx, Throwable throwable, StreamObserver streamObserver) { + ResponseWriter.write( + streamObserver, + ReceiveMessageResponse.newBuilder().setStatus(ResponseBuilder.buildStatus(throwable)).build() + ); + return null; } public CompletableFuture ackMessage(Context ctx, AckMessageRequest request) { @@ -375,12 +373,12 @@ public class ConsumerService extends BaseService { this.receiveMessageResultFilter = receiveMessageResultFilter; } - public ResponseHook> getReceiveMessageHook() { + public ResponseHook getReceiveMessageHook() { return receiveMessageHook; } public void setReceiveMessageHook( - ResponseHook> receiveMessageHook) { + ResponseHook receiveMessageHook) { this.receiveMessageHook = receiveMessageHook; } diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ReceiveMessageResponseStreamObserver.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ReceiveMessageResponseStreamObserver.java new file mode 100644 index 0000000000..d6946efb1b --- /dev/null +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ReceiveMessageResponseStreamObserver.java @@ -0,0 +1,61 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You 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.apache.rocketmq.proxy.grpc.v2.service.cluster; + +import apache.rocketmq.v2.ReceiveMessageRequest; +import apache.rocketmq.v2.ReceiveMessageResponse; +import io.grpc.Context; +import io.grpc.stub.StreamObserver; +import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseHook; + +public class ReceiveMessageResponseStreamObserver implements StreamObserver { + + private final Context context; + private final ReceiveMessageRequest request; + private final ResponseHook receiveMessageHook; + private final StreamObserver observer; + + public ReceiveMessageResponseStreamObserver(Context context, ReceiveMessageRequest request, + ResponseHook receiveMessageHook, + StreamObserver observer) { + this.context = context; + this.request = request; + this.receiveMessageHook = receiveMessageHook; + this.observer = observer; + } + + @Override + public void onNext(ReceiveMessageResponse response) { + if (receiveMessageHook != null) { + receiveMessageHook.beforeResponse(context, request, response, null); + } + observer.onNext(response); + } + + @Override + public void onError(Throwable throwable) { + if (receiveMessageHook != null) { + receiveMessageHook.beforeResponse(context, request, null, throwable); + } + observer.onError(throwable); + } + + @Override + public void onCompleted() { + observer.onCompleted(); + } +} \ No newline at end of file diff --git a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerServiceTest.java b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerServiceTest.java index 353c743bc1..35ed62dd26 100644 --- a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerServiceTest.java +++ b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/ConsumerServiceTest.java @@ -16,12 +16,14 @@ import apache.rocketmq.v2.RetryPolicy; import apache.rocketmq.v2.Settings; import apache.rocketmq.v2.Subscription; import io.grpc.Context; -import java.util.ArrayList; +import io.grpc.stub.StreamObserver; import java.util.List; +import java.util.Set; import java.util.concurrent.CompletableFuture; import java.util.concurrent.Executors; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicReference; +import java.util.stream.Collectors; import org.apache.rocketmq.client.consumer.AckResult; import org.apache.rocketmq.client.consumer.AckStatus; import org.apache.rocketmq.client.consumer.PopResult; @@ -36,19 +38,24 @@ import org.apache.rocketmq.proxy.connector.route.SelectableMessageQueue; import org.apache.rocketmq.remoting.protocol.RemotingCommand; import org.assertj.core.util.Lists; import org.junit.Test; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyLong; import static org.mockito.ArgumentMatchers.anyString; -import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; public class ConsumerServiceTest extends BaseServiceTest { @Mock private ReadQueueSelector readQueueSelector; + @Mock + private StreamObserver receiveMessageResponseStreamObserver; private ConsumerService consumerService; private DefaultReceiveMessageResultFilter receiveMessageResultFilter; @@ -90,7 +97,7 @@ public class ConsumerServiceTest extends BaseServiceTest { Context ctx = Context.current().withDeadlineAfter(3, TimeUnit.SECONDS, Executors.newSingleThreadScheduledExecutor()); AtomicReference ackHandler = new AtomicReference<>(); receiveMessageResultFilter.setAckNoMatchedMessageHook((ctx1, request, response, t) -> ackHandler.set(request.getExtraInfo())); - List responseList = consumerService.receiveMessage(ctx, + consumerService.receiveMessage(ctx, ReceiveMessageRequest.newBuilder() .setMessageQueue(apache.rocketmq.v2.MessageQueue.newBuilder() .setTopic(Resource.newBuilder() @@ -102,12 +109,19 @@ public class ConsumerServiceTest extends BaseServiceTest { .setType(FilterType.TAG) .setExpression("msg1") .build()) - .build() - ).get(); + .build(), + receiveMessageResponseStreamObserver + ); + ArgumentCaptor argument = ArgumentCaptor.forClass(ReceiveMessageResponse.class); + verify(receiveMessageResponseStreamObserver, times(2)).onNext(argument.capture()); + verify(receiveMessageResponseStreamObserver, times(1)).onCompleted(); - assertEquals(1, responseList.size()); - ReceiveMessageResponse response = responseList.get(0); + ReceiveMessageResponse response = argument.getAllValues().get(0); + assertTrue(response.hasStatus()); assertEquals(Code.OK, response.getStatus().getCode()); + + response = argument.getAllValues().get(1); + assertTrue(response.hasMessage()); assertEquals("msg1", response.getMessage().getSystemProperties().getMessageId()); assertEquals(ReceiptHandle.create(messageExtList.get(1)).getReceiptHandle(), ackHandler.get()); } @@ -135,15 +149,13 @@ public class ConsumerServiceTest extends BaseServiceTest { when(readConsumerClient.popMessage(any(), anyString(), anyString(), any(), anyLong())) .thenReturn(CompletableFuture.completedFuture(popResult)); when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr"); - List toDLQMsgId = new ArrayList<>(); - doAnswer(mock -> { - ConsumerSendMsgBackRequestHeader sendMsgBackRequestHeader = mock.getArgument(2); - toDLQMsgId.add(sendMsgBackRequestHeader.getOriginMsgId()); - return CompletableFuture.completedFuture(RemotingCommand.createResponseCommand(ResponseCode.SUCCESS, "")); - }).when(producerClient).sendMessageBackThenAckOrg(any(), anyString(), any(), any()); + ArgumentCaptor sendMsgBackRequestHeaderArgumentCaptor = + ArgumentCaptor.forClass(ConsumerSendMsgBackRequestHeader.class); + when(producerClient.sendMessageBackThenAckOrg(any(), anyString(), sendMsgBackRequestHeaderArgumentCaptor.capture(), any())) + .thenReturn(CompletableFuture.completedFuture(RemotingCommand.createResponseCommand(ResponseCode.SUCCESS, ""))); Context ctx = Context.current().withDeadlineAfter(3, TimeUnit.SECONDS, Executors.newSingleThreadScheduledExecutor()); - List responseList = consumerService.receiveMessage(ctx, + consumerService.receiveMessage(ctx, ReceiveMessageRequest.newBuilder() .setMessageQueue(apache.rocketmq.v2.MessageQueue.newBuilder() .setTopic(Resource.newBuilder() @@ -155,13 +167,20 @@ public class ConsumerServiceTest extends BaseServiceTest { .setType(FilterType.TAG) .setExpression("msg1") .build()) - .build() - ).get(); + .build(), + receiveMessageResponseStreamObserver + ); + ArgumentCaptor argument = ArgumentCaptor.forClass(ReceiveMessageResponse.class); + verify(receiveMessageResponseStreamObserver, times(1)).onNext(argument.capture()); + verify(receiveMessageResponseStreamObserver, times(1)).onCompleted(); - assertEquals(1, responseList.size()); - ReceiveMessageResponse response = responseList.get(0); + ReceiveMessageResponse response = argument.getValue(); assertEquals(Code.OK, response.getStatus().getCode()); - assertEquals(2, toDLQMsgId.size()); + assertEquals(2, sendMsgBackRequestHeaderArgumentCaptor.getAllValues().size()); + Set toDLQMsgId = sendMsgBackRequestHeaderArgumentCaptor.getAllValues().stream() + .map(ConsumerSendMsgBackRequestHeader::getOriginMsgId).collect(Collectors.toSet()); + assertTrue(toDLQMsgId.contains("msg1")); + assertTrue(toDLQMsgId.contains("msg2")); } @Test @@ -190,11 +209,9 @@ public class ConsumerServiceTest extends BaseServiceTest { @Test public void testNackMessageToDLQ() throws Exception { ReceiptHandle receiptHandle = createReceiptHandle(); - AtomicReference headerRef = new AtomicReference<>(); - doAnswer(mock -> { - headerRef.set(mock.getArgument(2)); - return CompletableFuture.completedFuture(RemotingCommand.createResponseCommand(ResponseCode.SUCCESS, "")); - }).when(producerClient).sendMessageBackThenAckOrg(any(), anyString(), any(), any()); + ArgumentCaptor headerArgumentCaptor = ArgumentCaptor.forClass(ConsumerSendMsgBackRequestHeader.class); + when(producerClient.sendMessageBackThenAckOrg(any(), anyString(), headerArgumentCaptor.capture(), any())) + .thenReturn(CompletableFuture.completedFuture(RemotingCommand.createResponseCommand(ResponseCode.SUCCESS, ""))); when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr"); Settings clientSettings = createClientSettings(3); @@ -213,19 +230,17 @@ public class ConsumerServiceTest extends BaseServiceTest { .get(); assertEquals(Code.OK, response.getStatus().getCode()); - assertEquals(receiptHandle.getCommitLogOffset(), headerRef.get().getOffset().longValue()); + assertEquals(receiptHandle.getCommitLogOffset(), headerArgumentCaptor.getValue().getOffset().longValue()); } @Test public void testNackMessage() throws Exception { ReceiptHandle receiptHandle = createReceiptHandle(); - AtomicReference headerRef = new AtomicReference<>(); - doAnswer(mock -> { - headerRef.set(mock.getArgument(3)); - AckResult ackResult = new AckResult(); - ackResult.setStatus(AckStatus.OK); - return CompletableFuture.completedFuture(ackResult); - }).when(writeConsumerClient).changeInvisibleTimeAsync(any(), anyString(), anyString(), any()); + ArgumentCaptor headerArgumentCaptor = ArgumentCaptor.forClass(ChangeInvisibleTimeRequestHeader.class); + AckResult ackResult = new AckResult(); + ackResult.setStatus(AckStatus.OK); + when(writeConsumerClient.changeInvisibleTimeAsync(any(), anyString(), anyString(), headerArgumentCaptor.capture())) + .thenReturn(CompletableFuture.completedFuture(ackResult)); when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr"); Settings clientSettings = createClientSettings(3); @@ -244,8 +259,8 @@ public class ConsumerServiceTest extends BaseServiceTest { .get(); assertEquals(Code.OK, response.getStatus().getCode()); - assertEquals(receiptHandle.getOffset(), headerRef.get().getOffset().longValue()); - assertEquals(receiptHandle.encode(), headerRef.get().getExtraInfo()); + assertEquals(receiptHandle.getOffset(), headerArgumentCaptor.getValue().getOffset().longValue()); + assertEquals(receiptHandle.encode(), headerArgumentCaptor.getValue().getExtraInfo()); } private Settings createClientSettings(int maxDeliveryAttempts) { diff --git a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/TransactionServiceTest.java b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/TransactionServiceTest.java index e9344d62c7..cccfe03ee4 100644 --- a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/TransactionServiceTest.java +++ b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/TransactionServiceTest.java @@ -5,7 +5,6 @@ import apache.rocketmq.v2.EndTransactionRequest; import apache.rocketmq.v2.EndTransactionResponse; import apache.rocketmq.v2.TelemetryCommand; import io.grpc.Context; -import java.util.concurrent.atomic.AtomicReference; import org.apache.rocketmq.common.protocol.header.EndTransactionRequestHeader; import org.apache.rocketmq.proxy.channel.ChannelManager; import org.apache.rocketmq.proxy.connector.transaction.TransactionId; @@ -14,13 +13,14 @@ import org.apache.rocketmq.proxy.grpc.v2.adapter.channel.GrpcClientChannel; import org.apache.rocketmq.remoting.common.RemotingHelper; import org.assertj.core.util.Lists; import org.junit.Test; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyString; -import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.doNothing; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; @@ -38,14 +38,11 @@ public class TransactionServiceTest extends BaseServiceTest { @Test public void testCheckTransactionState() { GrpcClientChannel channel = mock(GrpcClientChannel.class); - AtomicReference writeDataRef = new AtomicReference<>(); when(channelManager.getClientIdList(anyString())).thenReturn(Lists.newArrayList("clientId")); when(channelManager.getChannel(anyString(), any())).thenReturn(channel); - doAnswer(mock -> { - writeDataRef.set(mock.getArgument(0)); - return null; - }).when(channel).writeAndFlush(any()); + ArgumentCaptor flushDataCaptor = ArgumentCaptor.forClass(Object.class); + when(channel.writeAndFlush(flushDataCaptor.capture())).thenReturn(null); TransactionId transactionId = TransactionId.genByBrokerTransactionId( RemotingHelper.string2SocketAddress("127.0.0.1:8080"), @@ -59,23 +56,21 @@ public class TransactionServiceTest extends BaseServiceTest { createMessageExt("msgId", "msgId") )); - assertTrue(writeDataRef.get() instanceof TelemetryCommand); - TelemetryCommand response = (TelemetryCommand) writeDataRef.get(); + Object flushData = flushDataCaptor.getValue(); + assertTrue(flushData instanceof TelemetryCommand); + TelemetryCommand response = (TelemetryCommand) flushData; assertEquals(transactionId.getProxyTransactionId(), response.getRecoverOrphanedTransactionCommand().getTransactionId()); } @Test public void testEndTransaction() throws Exception { - AtomicReference headerRef = new AtomicReference<>(); - AtomicReference brokerAddrRef = new AtomicReference<>(); TransactionId transactionId = TransactionId.genByBrokerTransactionId( RemotingHelper.string2SocketAddress("127.0.0.1:8080"), "71F99B78B6E261357FA259CCA6456118", 1234, 5678); - doAnswer(mock -> { - brokerAddrRef.set(mock.getArgument(1)); - headerRef.set(mock.getArgument(2)); - return null; - }).when(producerClient).endTransaction(any(), anyString(), any()); + ArgumentCaptor brokerAddrCaptor = ArgumentCaptor.forClass(String.class); + ArgumentCaptor headerCaptor = ArgumentCaptor.forClass(EndTransactionRequestHeader.class); + doNothing().when(producerClient) + .endTransaction(any(), brokerAddrCaptor.capture(), headerCaptor.capture()); EndTransactionResponse response = transactionService.endTransaction(Context.current(), EndTransactionRequest.newBuilder() .setTransactionId(transactionId.getProxyTransactionId()) @@ -83,7 +78,7 @@ public class TransactionServiceTest extends BaseServiceTest { ).get(); assertEquals(Code.OK, response.getStatus().getCode()); - assertEquals(transactionId.getBrokerTransactionId(), headerRef.get().getTransactionId()); - assertEquals("127.0.0.1:8080", brokerAddrRef.get()); + assertEquals(transactionId.getBrokerTransactionId(), headerCaptor.getValue().getTransactionId()); + assertEquals("127.0.0.1:8080", brokerAddrCaptor.getValue()); } } \ No newline at end of file diff --git a/test/src/test/java/org/apache/rocketmq/test/grpc/v2/GrpcBaseTest.java b/test/src/test/java/org/apache/rocketmq/test/grpc/v2/GrpcBaseTest.java index 0ba85b8109..41c9d84426 100644 --- a/test/src/test/java/org/apache/rocketmq/test/grpc/v2/GrpcBaseTest.java +++ b/test/src/test/java/org/apache/rocketmq/test/grpc/v2/GrpcBaseTest.java @@ -223,9 +223,8 @@ public class GrpcBaseTest extends BaseConf { this.sendClientSettings(stub, buildPushConsumerClientSettings()).get(); - ReceiveMessageResponse response = receiveMessage(blockingStub, topic, group).get(0); - assertReceiveMessage(response, messageId); - String receiptHandle = response.getMessage().getSystemProperties().getReceiptHandle(); + Message responseMessage = assertAndGetReceiveMessage(receiveMessage(blockingStub, topic, group), messageId); + String receiptHandle = responseMessage.getSystemProperties().getReceiptHandle(); AckMessageResponse ackMessageResponse = blockingStub.ackMessage(buildAckMessageRequest(topic, group, messageId, receiptHandle)); assertAllAckOk(ackMessageResponse); } @@ -245,27 +244,26 @@ public class GrpcBaseTest extends BaseConf { this.sendClientSettings(stub, buildPushConsumerClientSettings()).get(); - ReceiveMessageResponse receiveResponse = receiveMessage(blockingStub, topic, group).get(0); - assertReceiveMessage(receiveResponse, messageId); + Message message = assertAndGetReceiveMessage(receiveMessage(blockingStub, topic, group), messageId); - Message message = receiveResponse.getMessage(); NackMessageResponse nackMessageResponse = blockingStub.nackMessage(buildNackMessageRequest( topic, group, messageId, message.getSystemProperties().getReceiptHandle(), 1 )); assertNackMessageResponse(nackMessageResponse); - AtomicReference receiveRetryResponseRef = new AtomicReference<>(); + AtomicReference receiveRetryMessageRef = new AtomicReference<>(); await().atMost(java.time.Duration.ofSeconds(30)).until(() -> { - ReceiveMessageResponse receiveRetryResponse = receiveMessage(blockingStub, topic, group, 1).get(0); - if (!receiveRetryResponse.hasMessage()) { + List messageList = getMessageFromReceiveMessageResponse(receiveMessage(blockingStub, topic, group, 1)); + if (messageList.isEmpty()) { return false; } - receiveRetryResponseRef.set(receiveRetryResponse); - return receiveRetryResponse.getMessage().getSystemProperties() + + receiveRetryMessageRef.set(messageList.get(0)); + return messageList.get(0).getSystemProperties() .getMessageId().equals(messageId); }); - message = receiveRetryResponseRef.get().getMessage(); + message = receiveRetryMessageRef.get(); nackMessageResponse = blockingStub.nackMessage(buildNackMessageRequest( topic, group, messageId, message.getSystemProperties().getReceiptHandle(), 2 )); @@ -361,11 +359,11 @@ public class GrpcBaseTest extends BaseConf { .build()); await().atMost(java.time.Duration.ofSeconds(30)).until(() -> { - ReceiveMessageResponse receiveRetryResponse = receiveMessage(blockingStub, topic, group).get(0); - if (!receiveRetryResponse.hasMessage()) { + List retryMessageList = getMessageFromReceiveMessageResponse(receiveMessage(blockingStub, topic, group)); + if (retryMessageList.isEmpty()) { return false; } - return receiveRetryResponse.getMessage().getSystemProperties() + return retryMessageList.get(0).getSystemProperties() .getMessageId().equals(messageId); }); } finally { @@ -390,10 +388,9 @@ public class GrpcBaseTest extends BaseConf { this.sendClientSettings(stub, buildSimpleConsumerClientSettings(maxDeliveryAttempts, fifo)).get(); - ReceiveMessageResponse receiveResponse = receiveMessage(blockingStub, topic, group).get(0); - assertReceiveMessage(receiveResponse, messageId); + Message message = assertAndGetReceiveMessage(receiveMessage(blockingStub, topic, group), messageId); - String receiptHandle = receiveResponse.getMessage().getSystemProperties().getReceiptHandle(); + String receiptHandle = message.getSystemProperties().getReceiptHandle(); ChangeInvisibleDurationResponse changeResponse = blockingStub.changeInvisibleDuration(buildChangeInvisibleDurationRequest(topic, group, receiptHandle, 5)); assertChangeInvisibleDurationResponse(changeResponse, receiptHandle); @@ -401,13 +398,13 @@ public class GrpcBaseTest extends BaseConf { ackHandles.add(changeResponse.getReceiptHandle()); await().atMost(java.time.Duration.ofSeconds(20)).until(() -> { - ReceiveMessageResponse receiveRetryResponse = receiveMessage(blockingStub, topic, group).get(0); - if (!receiveRetryResponse.hasMessage()) { + List retryMessageList = getMessageFromReceiveMessageResponse(receiveMessage(blockingStub, topic, group)); + if (retryMessageList.isEmpty()) { return false; } - if (receiveRetryResponse.getMessage().getSystemProperties() + if (retryMessageList.get(0).getSystemProperties() .getMessageId().equals(messageId)) { - ackHandles.add(receiveRetryResponse.getMessage().getSystemProperties().getReceiptHandle()); + ackHandles.add(retryMessageList.get(0).getSystemProperties().getReceiptHandle()); return true; } return false; @@ -450,8 +447,7 @@ public class GrpcBaseTest extends BaseConf { AtomicInteger receiveMessageCount = new AtomicInteger(0); - ReceiveMessageResponse receiveResponse = receiveMessage(blockingStub, topic, group).get(0); - assertReceiveMessage(receiveResponse, messageId); + assertAndGetReceiveMessage(receiveMessage(blockingStub, topic, group), messageId); receiveMessageCount.incrementAndGet(); DefaultMQPullConsumer defaultMQPullConsumer = new DefaultMQPullConsumer(group); @@ -459,10 +455,8 @@ public class GrpcBaseTest extends BaseConf { org.apache.rocketmq.common.message.MessageQueue dlqMQ = new org.apache.rocketmq.common.message.MessageQueue(MixAll.getDLQTopic(group), broker1Name, 0); await().atMost(java.time.Duration.ofSeconds(30)).until(() -> { try { - ReceiveMessageResponse retryReceiveResponse = receiveMessage(blockingStub, topic, group, 1).get(0); - if (retryReceiveResponse.hasMessage()) { - receiveMessageCount.incrementAndGet(); - } + List messageList = getMessageFromReceiveMessageResponse(receiveMessage(blockingStub, topic, group, 1)); + receiveMessageCount.addAndGet(messageList.size()); PullResult pullResult = defaultMQPullConsumer.pull(dlqMQ, "*", 0L, 1); if (!PullStatus.FOUND.equals(pullResult.getPullStatus())) { @@ -480,13 +474,7 @@ public class GrpcBaseTest extends BaseConf { public List receiveMessage(MessagingServiceGrpc.MessagingServiceBlockingStub stub, String topic, String group) { - List responseList = new ArrayList<>(); - Iterator responseIterator = stub.withDeadlineAfter(15, TimeUnit.SECONDS) - .receiveMessage(buildReceiveMessageRequest(topic, group)); - while (responseIterator.hasNext()) { - responseList.add(responseIterator.next()); - } - return responseList; + return receiveMessage(stub, topic, group, 15); } public List receiveMessage(MessagingServiceGrpc.MessagingServiceBlockingStub stub, @@ -500,6 +488,16 @@ public class GrpcBaseTest extends BaseConf { return responseList; } + public List getMessageFromReceiveMessageResponse(List responseList) { + List messageList = new ArrayList<>(); + for (ReceiveMessageResponse response : responseList) { + if (response.hasMessage()) { + messageList.add(response.getMessage()); + } + } + return messageList; + } + public QueryRouteRequest buildQueryRouteRequest(String topic) { return QueryRouteRequest.newBuilder() .setTopic(Resource.newBuilder() @@ -647,12 +645,14 @@ public class GrpcBaseTest extends BaseConf { assertThat(response.getReceipts(0).getMessageId()).isEqualTo(messageId); } - public void assertReceiveMessage(ReceiveMessageResponse response, String messageId) { - assertThat(response.getStatus() + public Message assertAndGetReceiveMessage(List response, String messageId) { + assertThat(response.get(0).hasStatus()).isTrue(); + assertThat(response.get(0).getStatus() .getCode()).isEqualTo(Code.OK); - assertThat(response.getMessage() + assertThat(response.get(1).getMessage() .getSystemProperties() .getMessageId()).isEqualTo(messageId); + return response.get(1).getMessage(); } public void assertAllAckOk(AckMessageResponse response) {