diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/ReceiveMessageResponseStreamWriter.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/ReceiveMessageResponseStreamWriter.java index 8723558b71..32998280f2 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/ReceiveMessageResponseStreamWriter.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/ReceiveMessageResponseStreamWriter.java @@ -16,19 +16,25 @@ */ package org.apache.rocketmq.proxy.grpc.v2.service; +import apache.rocketmq.v2.Code; +import apache.rocketmq.v2.Message; import apache.rocketmq.v2.ReceiveMessageRequest; import apache.rocketmq.v2.ReceiveMessageResponse; import io.grpc.Context; import io.grpc.stub.StreamObserver; +import java.util.Iterator; import java.util.List; import org.apache.rocketmq.client.consumer.PopStatus; import org.apache.rocketmq.common.message.MessageExt; +import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseBuilder; import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseHook; +import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseWriter; public abstract class ReceiveMessageResponseStreamWriter { protected final StreamObserver streamObserver; protected final ResponseHook receiveMessageHook; + protected final ReceiveMessageResultFilter receiveMessageResultFilter; public interface Builder { ReceiveMessageResponseStreamWriter build( @@ -38,12 +44,79 @@ public abstract class ReceiveMessageResponseStreamWriter { public ReceiveMessageResponseStreamWriter( StreamObserver observer, - ResponseHook hook) { + ResponseHook hook, + ReceiveMessageResultFilter messageResultFilter) { streamObserver = observer; receiveMessageHook = hook; + receiveMessageResultFilter = messageResultFilter; } - public abstract void write(Context ctx, ReceiveMessageRequest request, PopStatus status, List messageFoundList); + public void write(Context ctx, ReceiveMessageRequest request, PopStatus status, List messageFoundList) { + ReceiveMessageResponseStreamObserver responseStreamObserver = new ReceiveMessageResponseStreamObserver( + ctx, + request, + receiveMessageHook, + streamObserver); + try { + switch (status) { + case FOUND: + List messageList = this.receiveMessageResultFilter.filterMessage(ctx, request, messageFoundList); + if (messageList.isEmpty()) { + responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) + .build()); + } else { + responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.OK, Code.OK.name())) + .build()); + Iterator messageIterator = messageList.iterator(); + while (messageIterator.hasNext()) { + Message curMessage = messageIterator.next(); + try { + responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setMessage(curMessage) + .build()); + } catch (Throwable t) { + this.processThrowableWhenWriteMessage(t, ctx, request, curMessage); + messageIterator.forEachRemaining(message -> + this.processThrowableWhenWriteMessage(t, ctx, request, message)); + return; + } + } + } + break; + case POLLING_FULL: + responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.TOO_MANY_REQUESTS, "polling full")) + .build()); + break; + case NO_NEW_MSG: + case POLLING_NOT_FOUND: + default: + responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() + .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) + .build()); + break; + } + } catch (Throwable t) { + write(ctx, request, t); + } finally { + responseStreamObserver.onCompleted(); + } + } - public abstract void write(Context ctx, ReceiveMessageRequest request, Throwable throwable); + protected abstract void processThrowableWhenWriteMessage(Throwable throwable, + Context context, ReceiveMessageRequest request, Message message); + + public void write(Context ctx, ReceiveMessageRequest request, Throwable throwable) { + ReceiveMessageResponseStreamObserver responseStreamObserver = new ReceiveMessageResponseStreamObserver( + ctx, + request, + receiveMessageHook, + streamObserver); + ResponseWriter.write( + responseStreamObserver, + ReceiveMessageResponse.newBuilder().setStatus(ResponseBuilder.buildStatus(throwable)).build() + ); + } } diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/DefaultReceiveMessageResponseStreamWriter.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/DefaultReceiveMessageResponseStreamWriter.java index 0c8f50547a..bf071719f3 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/DefaultReceiveMessageResponseStreamWriter.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/cluster/DefaultReceiveMessageResponseStreamWriter.java @@ -16,31 +16,23 @@ */ package org.apache.rocketmq.proxy.grpc.v2.service.cluster; -import apache.rocketmq.v2.Code; import apache.rocketmq.v2.Message; import apache.rocketmq.v2.ReceiveMessageRequest; import apache.rocketmq.v2.ReceiveMessageResponse; import io.grpc.Context; import io.grpc.stub.StreamObserver; import java.time.Duration; -import java.util.Iterator; -import java.util.List; import org.apache.rocketmq.client.consumer.AckStatus; -import org.apache.rocketmq.client.consumer.PopStatus; import org.apache.rocketmq.common.constant.LoggerName; import org.apache.rocketmq.common.consumer.ReceiptHandle; -import org.apache.rocketmq.common.message.MessageExt; import org.apache.rocketmq.common.protocol.header.ChangeInvisibleTimeRequestHeader; import org.apache.rocketmq.logging.InternalLogger; import org.apache.rocketmq.logging.InternalLoggerFactory; import org.apache.rocketmq.proxy.connector.ForwardWriteConsumer; import org.apache.rocketmq.proxy.connector.route.TopicRouteCache; import org.apache.rocketmq.proxy.grpc.v2.adapter.GrpcConverter; -import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseBuilder; import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseHook; -import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseWriter; import org.apache.rocketmq.proxy.grpc.v2.service.BaseService; -import org.apache.rocketmq.proxy.grpc.v2.service.ReceiveMessageResponseStreamObserver; import org.apache.rocketmq.proxy.grpc.v2.service.ReceiveMessageResponseStreamWriter; import org.apache.rocketmq.proxy.grpc.v2.service.ReceiveMessageResultFilter; @@ -50,7 +42,6 @@ public class DefaultReceiveMessageResponseStreamWriter extends ReceiveMessageRes protected static final long NACK_INVISIBLE_TIME = Duration.ofSeconds(1).toMillis(); protected final ForwardWriteConsumer writeConsumer; protected final TopicRouteCache topicRouteCache; - protected volatile ReceiveMessageResultFilter receiveMessageResultFilter; public DefaultReceiveMessageResponseStreamWriter( StreamObserver observer, @@ -58,74 +49,15 @@ public class DefaultReceiveMessageResponseStreamWriter extends ReceiveMessageRes ForwardWriteConsumer writeConsumer, TopicRouteCache topicRouteCache, ReceiveMessageResultFilter receiveMessageResultFilter) { - super(observer, hook); + super(observer, hook, receiveMessageResultFilter); this.writeConsumer = writeConsumer; this.topicRouteCache = topicRouteCache; - this.receiveMessageResultFilter = receiveMessageResultFilter; } @Override - public void write(Context ctx, ReceiveMessageRequest request, PopStatus status, List messageFoundList) { - ReceiveMessageResponseStreamObserver responseStreamObserver = new ReceiveMessageResponseStreamObserver( - ctx, - request, - receiveMessageHook, - streamObserver); - try { - switch (status) { - case FOUND: - List messageList = this.receiveMessageResultFilter.filterMessage(ctx, request, messageFoundList); - if (messageList.isEmpty()) { - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) - .build()); - } else { - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, Code.OK.name())) - .build()); - Iterator messageIterator = messageList.iterator(); - while (messageIterator.hasNext()) { - if (responseStreamObserver.isCancelled()) { - break; - } - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setMessage(messageIterator.next()) - .build()); - } - messageIterator.forEachRemaining(message -> this.nackFailToWriteMessage(ctx, request, message)); - } - break; - case POLLING_FULL: - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.TOO_MANY_REQUESTS, "polling full")) - .build()); - break; - case NO_NEW_MSG: - case POLLING_NOT_FOUND: - default: - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) - .build()); - break; - } - } catch (Throwable t) { - write(ctx, request, t); - } finally { - responseStreamObserver.onCompleted(); - } - } - - @Override - public void write(Context ctx, ReceiveMessageRequest request, Throwable throwable) { - ReceiveMessageResponseStreamObserver responseStreamObserver = new ReceiveMessageResponseStreamObserver( - ctx, - request, - receiveMessageHook, - streamObserver); - ResponseWriter.write( - responseStreamObserver, - ReceiveMessageResponse.newBuilder().setStatus(ResponseBuilder.buildStatus(throwable)).build() - ); + protected void processThrowableWhenWriteMessage(Throwable throwable, Context context, ReceiveMessageRequest request, + Message message) { + this.nackFailToWriteMessage(context, request, message); } protected void nackFailToWriteMessage(Context ctx, ReceiveMessageRequest request, Message message) { @@ -161,13 +93,4 @@ public class DefaultReceiveMessageResponseStreamWriter extends ReceiveMessageRes log.warn("err when nack message. request:{}, message:{}", request, message, t); } } - - public ReceiveMessageResultFilter getReceiveMessageResultFilter() { - return receiveMessageResultFilter; - } - - public void setReceiveMessageResultFilter( - ReceiveMessageResultFilter receiveMessageResultFilter) { - this.receiveMessageResultFilter = receiveMessageResultFilter; - } } diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriter.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriter.java index 258afd1058..50fe0e1a59 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriter.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriter.java @@ -17,29 +17,21 @@ package org.apache.rocketmq.proxy.grpc.v2.service.local; -import apache.rocketmq.v2.Code; import apache.rocketmq.v2.Message; import apache.rocketmq.v2.ReceiveMessageRequest; import apache.rocketmq.v2.ReceiveMessageResponse; import io.grpc.Context; import io.grpc.stub.StreamObserver; import io.netty.channel.Channel; -import java.util.Iterator; -import java.util.List; import org.apache.rocketmq.broker.BrokerController; -import org.apache.rocketmq.client.consumer.PopStatus; import org.apache.rocketmq.common.constant.LoggerName; import org.apache.rocketmq.common.consumer.ReceiptHandle; -import org.apache.rocketmq.common.message.MessageExt; import org.apache.rocketmq.common.protocol.RequestCode; import org.apache.rocketmq.common.protocol.header.ChangeInvisibleTimeRequestHeader; import org.apache.rocketmq.proxy.channel.ChannelManager; import org.apache.rocketmq.proxy.channel.SimpleChannelHandlerContext; import org.apache.rocketmq.proxy.grpc.v2.adapter.GrpcConverter; -import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseBuilder; import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseHook; -import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseWriter; -import org.apache.rocketmq.proxy.grpc.v2.service.ReceiveMessageResponseStreamObserver; import org.apache.rocketmq.proxy.grpc.v2.service.ReceiveMessageResponseStreamWriter; import org.apache.rocketmq.proxy.grpc.v2.service.ReceiveMessageResultFilter; import org.apache.rocketmq.remoting.exception.RemotingCommandException; @@ -51,7 +43,6 @@ public class LocalReceiveMessageResponseStreamWriter extends ReceiveMessageRespo private final static Logger log = LoggerFactory.getLogger(LoggerName.PROXY_LOGGER_NAME); private final ChannelManager channelManager; private final BrokerController brokerController; - private final ReceiveMessageResultFilter receiveMessageResultFilter; public LocalReceiveMessageResponseStreamWriter( StreamObserver observer, @@ -59,73 +50,15 @@ public class LocalReceiveMessageResponseStreamWriter extends ReceiveMessageRespo ChannelManager channelManager, BrokerController brokerController, ReceiveMessageResultFilter receiveMessageResultFilter) { - super(observer, hook); + super(observer, hook, receiveMessageResultFilter); this.channelManager = channelManager; this.brokerController = brokerController; - this.receiveMessageResultFilter = receiveMessageResultFilter; } @Override - public void write(Context ctx, ReceiveMessageRequest request, PopStatus status, List messageFoundList) { - ReceiveMessageResponseStreamObserver responseStreamObserver = new ReceiveMessageResponseStreamObserver( - ctx, - request, - receiveMessageHook, - streamObserver); - try { - switch (status) { - case FOUND: - List messageList = this.receiveMessageResultFilter.filterMessage(ctx, request, messageFoundList); - if (messageList.isEmpty()) { - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) - .build()); - } else { - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, Code.OK.name())) - .build()); - Iterator messageIterator = messageList.iterator(); - while (messageIterator.hasNext()) { - if (responseStreamObserver.isCancelled()) { - break; - } - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setMessage(messageIterator.next()) - .build()); - } - messageIterator.forEachRemaining(message -> this.changeInvisibleTime(ctx, request, ReceiptHandle.decode(message.getSystemProperties().getReceiptHandle()))); - } - break; - case POLLING_FULL: - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.TOO_MANY_REQUESTS, "polling full")) - .build()); - break; - case NO_NEW_MSG: - case POLLING_NOT_FOUND: - default: - responseStreamObserver.onNext(ReceiveMessageResponse.newBuilder() - .setStatus(ResponseBuilder.buildStatus(Code.OK, "no new message")) - .build()); - break; - } - } catch (Throwable t) { - write(ctx, request, t); - } finally { - responseStreamObserver.onCompleted(); - } - } - - @Override public void write(Context ctx, ReceiveMessageRequest request, Throwable throwable) { - ReceiveMessageResponseStreamObserver responseStreamObserver = new ReceiveMessageResponseStreamObserver( - ctx, - request, - receiveMessageHook, - streamObserver); - ResponseWriter.write( - responseStreamObserver, - ReceiveMessageResponse.newBuilder().setStatus(ResponseBuilder.buildStatus(throwable)).build() - ); + protected void processThrowableWhenWriteMessage(Throwable throwable, Context context, ReceiveMessageRequest request, + Message message) { + this.changeInvisibleTime(context, request, ReceiptHandle.decode(message.getSystemProperties().getReceiptHandle())); } private void changeInvisibleTime(Context ctx, ReceiveMessageRequest request, ReceiptHandle handle) { diff --git a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/LocalGrpcServiceTest.java b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/LocalGrpcServiceTest.java index 17a1451fd9..dc07124b75 100644 --- a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/LocalGrpcServiceTest.java +++ b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/LocalGrpcServiceTest.java @@ -91,6 +91,7 @@ import org.apache.rocketmq.proxy.connector.transaction.TransactionId; import org.apache.rocketmq.proxy.grpc.interceptor.InterceptorConstants; import org.apache.rocketmq.proxy.grpc.v2.adapter.GrpcConverter; import org.apache.rocketmq.proxy.grpc.v2.adapter.ResponseBuilder; +import org.apache.rocketmq.remoting.common.RemotingUtil; import org.apache.rocketmq.remoting.exception.RemotingCommandException; import org.apache.rocketmq.remoting.netty.NettyRemotingServer; import org.apache.rocketmq.remoting.protocol.RemotingCommand; @@ -338,7 +339,7 @@ public class LocalGrpcServiceTest extends InitConfigAndLoggerTest { .setSystemProperties( message.getSystemProperties() .toBuilder() - .setReceiptHandle("0 0 1000 0 0 zhouxiang_MBP16 0 0 0") + .setReceiptHandle("0 0 1000 0 0 "+ brokerControllerMock.getBrokerConfig().getBrokerName() +" 0 0 0") .build()) .build()) .build(); diff --git a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriterTest.java b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriterTest.java index 5a53f5861a..a560673503 100644 --- a/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriterTest.java +++ b/proxy/src/test/java/org/apache/rocketmq/proxy/grpc/v2/service/local/LocalReceiveMessageResponseStreamWriterTest.java @@ -22,11 +22,14 @@ import apache.rocketmq.v2.Message; import apache.rocketmq.v2.ReceiveMessageRequest; import apache.rocketmq.v2.ReceiveMessageResponse; import io.grpc.Context; +import io.grpc.Status; +import io.grpc.StatusRuntimeException; import io.grpc.stub.ServerCallStreamObserver; import java.net.InetSocketAddress; import java.nio.charset.StandardCharsets; import java.util.ArrayList; import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; import org.apache.rocketmq.broker.BrokerController; import org.apache.rocketmq.broker.processor.ChangeInvisibleTimeProcessor; import org.apache.rocketmq.client.consumer.PopStatus; @@ -45,6 +48,7 @@ import org.junit.runner.RunWith; import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.Mockito; +import org.mockito.invocation.InvocationOnMock; import org.mockito.junit.MockitoJUnitRunner; import org.mockito.stubbing.Answer; @@ -120,7 +124,14 @@ public class LocalReceiveMessageResponseStreamWriterTest { @Test public void testWriteWhenCancel() throws RemotingCommandException { - Mockito.when(streamObserverMock.isCancelled()).thenReturn(true); + AtomicInteger onNextCallTimes = new AtomicInteger(0); + Mockito.doAnswer(mock -> { + if (onNextCallTimes.get() <=0) { + onNextCallTimes.incrementAndGet(); + return null; + } + throw new StatusRuntimeException(Status.CANCELLED); + }).when(streamObserverMock).onNext(Mockito.any()); Mockito.when(brokerControllerMock.getChangeInvisibleTimeProcessor()).thenReturn(changeInvisibleTimeProcessorMock); MessageExt messageExt = new MessageExt(); messageExt.setTopic("topic");