[ISSUE #3949] add test cases

This commit is contained in:
kaiyi.lk
2022-07-13 11:29:13 +08:00
committed by zhouxiang
parent 36b1e9fa2a
commit a613f6c4f3
29 changed files with 985 additions and 160 deletions
@@ -66,6 +66,7 @@ public class ProxyStartup {
final HealthCheckServer healthCheckServer = new HealthCheckServer();
PROXY_START_AND_SHUTDOWN.appendStartAndShutdown(healthCheckServer);
PROXY_START_AND_SHUTDOWN.start();
Runtime.getRuntime().addShutdownHook(new Thread(() -> {
LOGGER.info("try to shutdown server");
try {
@@ -105,9 +105,13 @@ public class ChannelManager {
}
public static SimpleChannel createSimpleChannelDirectly() {
final String clientHost = InterceptorConstants.METADATA.get(Context.current())
return createSimpleChannelDirectly(Context.current());
}
public static SimpleChannel createSimpleChannelDirectly(Context ctx) {
final String clientHost = InterceptorConstants.METADATA.get(ctx)
.get(InterceptorConstants.REMOTE_ADDRESS);
final String localAddress = InterceptorConstants.METADATA.get(Context.current())
final String localAddress = InterceptorConstants.METADATA.get(ctx)
.get(InterceptorConstants.LOCAL_ADDRESS);
return new SimpleChannel(null, clientHost, localAddress, ConfigurationManager.getProxyConfig().getChannelExpiredInSeconds());
}
@@ -58,19 +58,7 @@ public class ForwardProducer extends AbstractForwardClient {
return this.getClient().sendHeartbeat(heartbeatAddr, heartbeatData, timeout);
}
public void endTransaction(EndTransactionRequestHeader request, long timeoutMillis) throws Exception {
TransactionId transactionId = TransactionId.decode(request.getTransactionId());
EndTransactionRequestHeader requestHeader = new EndTransactionRequestHeader();
requestHeader.setProducerGroup(request.getProducerGroup());
requestHeader.setTranStateTableOffset(transactionId.getTranStateTableOffset());
requestHeader.setCommitLogOffset(transactionId.getCommitLogOffset());
requestHeader.setFromTransactionCheck(request.getFromTransactionCheck());
requestHeader.setMsgId(request.getMsgId());
requestHeader.setTransactionId(transactionId.getBrokerTransactionId());
requestHeader.setCommitOrRollback(request.getCommitOrRollback());
String brokerAddr = RemotingHelper.parseSocketAddressAddr(transactionId.getBrokerAddr());
public void endTransaction(String brokerAddr, EndTransactionRequestHeader requestHeader, long timeoutMillis) throws Exception {
this.getClient().endTransactionOneway(
brokerAddr,
requestHeader,
@@ -50,13 +50,19 @@ import apache.rocketmq.v1.ReportMessageConsumptionResultRequest;
import apache.rocketmq.v1.ReportMessageConsumptionResultResponse;
import apache.rocketmq.v1.ReportThreadStackTraceRequest;
import apache.rocketmq.v1.ReportThreadStackTraceResponse;
import apache.rocketmq.v1.ResponseCommon;
import apache.rocketmq.v1.SendMessageRequest;
import apache.rocketmq.v1.SendMessageResponse;
import com.google.rpc.Code;
import io.grpc.Context;
import io.grpc.stub.StreamObserver;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.CompletionException;
import org.apache.rocketmq.common.constant.LoggerName;
import org.apache.rocketmq.proxy.grpc.adapter.ResponseWriter;
import org.apache.rocketmq.proxy.grpc.common.ProxyException;
import org.apache.rocketmq.proxy.grpc.common.ResponseBuilder;
import org.apache.rocketmq.proxy.grpc.common.ResponseWriter;
import org.apache.rocketmq.proxy.grpc.service.GrpcForwardService;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
@@ -69,12 +75,25 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
this.grpcForwardService = grpcForwardService;
}
public ResponseCommon convertExceptionToResponseCommon(Throwable t) {
if (t instanceof CompletionException) {
if (t.getCause() instanceof ProxyException) {
ProxyException proxyException = (ProxyException) t.getCause();
return ResponseBuilder.buildCommon(proxyException.getCode(), proxyException.getMessage());
}
}
return ResponseBuilder.buildCommon(Code.INTERNAL, "internal error");
}
@Override
public void queryRoute(QueryRouteRequest request, StreamObserver<QueryRouteResponse> responseObserver) {
CompletableFuture<QueryRouteResponse> future = grpcForwardService.queryRoute(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
QueryRouteResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -84,7 +103,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<HeartbeatResponse> future = grpcForwardService.heartbeat(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
HeartbeatResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -94,7 +116,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<HealthCheckResponse> future = grpcForwardService.healthCheck(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
HealthCheckResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -104,7 +129,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<SendMessageResponse> future = grpcForwardService.sendMessage(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
SendMessageResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -114,7 +142,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<QueryAssignmentResponse> future = grpcForwardService.queryAssignment(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
QueryAssignmentResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -124,7 +155,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<ReceiveMessageResponse> future = grpcForwardService.receiveMessage(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
ReceiveMessageResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -134,7 +168,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<AckMessageResponse> future = grpcForwardService.ackMessage(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
AckMessageResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -144,7 +181,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<NackMessageResponse> future = grpcForwardService.nackMessage(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
NackMessageResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -154,7 +194,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<ForwardMessageToDeadLetterQueueResponse> future = grpcForwardService.forwardMessageToDeadLetterQueue(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
ForwardMessageToDeadLetterQueueResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -164,7 +207,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<EndTransactionResponse> future = grpcForwardService.endTransaction(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
EndTransactionResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -174,7 +220,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<QueryOffsetResponse> future = grpcForwardService.queryOffset(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
QueryOffsetResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -184,7 +233,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<PullMessageResponse> future = grpcForwardService.pullMessage(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
PullMessageResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -204,7 +256,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<ReportThreadStackTraceResponse> future = grpcForwardService.reportThreadStackTrace(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
ReportThreadStackTraceResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -214,7 +269,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<ReportMessageConsumptionResultResponse> future = grpcForwardService.reportMessageConsumptionResult(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
ReportMessageConsumptionResultResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -224,7 +282,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<NotifyClientTerminationResponse> future = grpcForwardService.notifyClientTermination(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
NotifyClientTerminationResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -234,7 +295,10 @@ public class GrpcMessagingProcessor extends MessagingServiceGrpc.MessagingServic
CompletableFuture<ChangeInvisibleDurationResponse> future = grpcForwardService.changeInvisibleDuration(Context.current(), request);
future.thenAccept(response -> ResponseWriter.write(responseObserver, response))
.exceptionally(e -> {
ResponseWriter.writeException(responseObserver, e);
ResponseWriter.write(
responseObserver,
ChangeInvisibleDurationResponse.newBuilder().setCommon(convertExceptionToResponseCommon(e)).build()
);
return null;
});
}
@@ -51,11 +51,12 @@ public class DelayPolicy {
private static List<Long> buildList(String messageDelayLevel) {
List<String> delayLevelList = Lists.newArrayList(Splitter.on(" ").split(messageDelayLevel));
List<Long> delayIntervalList = new ArrayList<>();
// the index of messageDelayLevel start from 1, so add a default value
delayIntervalList.add(0L);
for (String delayLevel : delayLevelList) {
final Pattern p = Pattern.compile("(\\d+)([smhd])");
final Matcher m = p.matcher(delayLevel);
while (m.find())
{
while (m.find()) {
final int duration = Integer.parseInt(m.group(1));
final String timeUnitString = m.group(2);
final long interval = toInterval(duration, timeUnitString);
@@ -263,7 +263,7 @@ public class GrpcConverter {
try {
handle = TransactionId.decode(transactionId);
} catch (UnknownHostException e) {
throw new IllegalArgumentException("Parse transaction id failed", e);
throw new ProxyException(Code.INVALID_ARGUMENT, "Parse transaction id failed", e);
}
long transactionStateTableOffset = handle.getTranStateTableOffset();
long commitLogOffset = handle.getCommitLogOffset();
@@ -273,7 +273,7 @@ public class GrpcConverter {
EndTransactionRequestHeader endTransactionRequestHeader = new EndTransactionRequestHeader();
endTransactionRequestHeader.setProducerGroup(groupName);
endTransactionRequestHeader.setMsgId(messageId);
endTransactionRequestHeader.setTransactionId(transactionId);
endTransactionRequestHeader.setTransactionId(handle.getBrokerTransactionId());
endTransactionRequestHeader.setTranStateTableOffset(transactionStateTableOffset);
endTransactionRequestHeader.setCommitLogOffset(commitLogOffset);
endTransactionRequestHeader.setCommitOrRollback(commitOrRollback);
@@ -313,7 +313,7 @@ public class GrpcConverter {
Map<String, String> userProperties = message.getUserAttributeMap();
for (String key : userProperties.keySet()) {
if (MessageConst.STRING_HASH_SET.contains(key)) {
throw new IllegalArgumentException("Property is used by system: " + key);
throw new ProxyException(Code.INVALID_ARGUMENT, "property is used by system: " + key);
}
}
MessageAccessor.setProperties(messageWithHeader, Maps.newHashMap(userProperties));
@@ -333,7 +333,7 @@ public class GrpcConverter {
// set message id
String messageId = message.getSystemAttribute().getMessageId();
if ("".equals(messageId)) {
throw new IllegalArgumentException("message id is empty");
throw new ProxyException(Code.INVALID_ARGUMENT, "message id is empty");
}
MessageAccessor.putProperty(messageWithHeader, MessageConst.PROPERTY_UNIQ_CLIENT_MESSAGE_ID_KEYIDX, messageId);
@@ -363,7 +363,7 @@ public class GrpcConverter {
case TIMEDDELIVERY_NOT_SET:
break;
default:
throw new IllegalStateException("Unexpected value: " + message.getSystemAttribute().getTimedDeliveryCase());
throw new ProxyException(Code.INVALID_ARGUMENT, "unexpected value: " + message.getSystemAttribute().getTimedDeliveryCase());
}
// set reconsume times
int reconsumeTimes = message.getSystemAttribute().getDeliveryAttempt();
@@ -16,7 +16,9 @@
*/
package org.apache.rocketmq.proxy.grpc.adapter;
import io.grpc.Context;
public interface ResponseHook<T, R> {
void beforeResponse(T request, R response, Throwable t);
void beforeResponse(Context ctx, T request, R response, Throwable t);
}
@@ -44,9 +44,9 @@ public class ResponseWriter {
}
}
public static <T> void writeException(StreamObserver<T> observer, final Throwable e) {
public static <T> void writeException(StreamObserver<?> observer, final Throwable e) {
if (observer instanceof ServerCallStreamObserver) {
final ServerCallStreamObserver<T> serverCallStreamObserver = (ServerCallStreamObserver<T>) observer;
final ServerCallStreamObserver serverCallStreamObserver = (ServerCallStreamObserver<T>) observer;
if (null == e) {
return;
}
@@ -56,6 +56,15 @@ public class ResponseWriter {
return;
}
// if (e instanceof CompletionException) {
// if (e.getCause() instanceof ProxyException) {
// ProxyException proxyException = (ProxyException) e.getCause();
// serverCallStreamObserver.onNext(ResponseBuilder.buildCommon(proxyException.getCode(), proxyException.getMessage()));
// serverCallStreamObserver.onCompleted();
// return;
// }
// }
LOGGER.debug("Start to write error response", e);
serverCallStreamObserver.onError(e);
serverCallStreamObserver.onCompleted();
@@ -19,6 +19,7 @@ package org.apache.rocketmq.proxy.grpc.adapter.channel;
import apache.rocketmq.v1.PollCommandResponse;
import apache.rocketmq.v1.PrintThreadStackTraceCommand;
import apache.rocketmq.v1.RecoverOrphanedTransactionCommand;
import io.grpc.Context;
import io.netty.channel.ChannelFuture;
import java.nio.ByteBuffer;
import java.util.concurrent.CompletableFuture;
@@ -42,7 +43,11 @@ public class GrpcClientChannel extends SimpleChannel {
private final PollResponseManager manager;
private GrpcClientChannel(String group, String clientId, PollResponseManager manager) {
super(ChannelManager.createSimpleChannelDirectly());
this(Context.current(), group, clientId, manager);
}
private GrpcClientChannel(Context ctx, String group, String clientId, PollResponseManager manager) {
super(ChannelManager.createSimpleChannelDirectly(ctx));
this.group = group;
this.clientId = clientId;
this.manager = manager;
@@ -57,10 +62,20 @@ public class GrpcClientChannel extends SimpleChannel {
String group,
String clientId,
PollResponseManager manager
) {
return create(Context.current(), channelManager, group, clientId, manager);
}
public static GrpcClientChannel create(
Context ctx,
ChannelManager channelManager,
String group,
String clientId,
PollCommandResponseManager manager
) {
GrpcClientChannel channel = channelManager.createChannel(
buildKey(group, clientId),
() -> new GrpcClientChannel(group, clientId, manager),
() -> new GrpcClientChannel(ctx, group, clientId, manager),
GrpcClientChannel.class
);
@@ -80,6 +95,14 @@ public class GrpcClientChannel extends SimpleChannel {
return group + "@" + clientId;
}
@Override
public boolean isWritable() {
if (this.pollCommandResponseFutureRef.get() == null) {
return false;
}
return !this.pollCommandResponseFutureRef.get().isDone();
}
/**
* Write response to corresponding remote client
*
@@ -114,7 +114,7 @@ public class ClusterGrpcService extends AbstractStartAndShutdown implements Grpc
@Override
public CompletableFuture<HeartbeatResponse> heartbeat(Context ctx, HeartbeatRequest request) {
this.clientService.heartbeat(ctx, request, channelManager);
this.clientService.heartbeat(ctx, request);
return CompletableFuture.completedFuture(
HeartbeatResponse.newBuilder()
.setCommon(ResponseBuilder.buildCommon(Code.OK, Code.OK.name()))
@@ -196,8 +196,8 @@ public class ClusterGrpcService extends AbstractStartAndShutdown implements Grpc
@Override
public CompletableFuture<NotifyClientTerminationResponse> notifyClientTermination(Context ctx,
NotifyClientTerminationRequest request) {
this.clientService.unregister(ctx, request, channelManager);
NotifyClientTerminationRequest request) {
this.clientService.unregister(ctx, request);
return CompletableFuture.completedFuture(
NotifyClientTerminationResponse.newBuilder()
.setCommon(ResponseBuilder.buildCommon(Code.OK, Code.OK.name()))
@@ -83,13 +83,17 @@ public class ConsumerService extends BaseService {
CompletableFuture<ReceiveMessageResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (receiveMessageHook != null) {
receiveMessageHook.beforeResponse(request, response, throwable);
receiveMessageHook.beforeResponse(ctx, request, response, throwable);
}
});
try {
PopMessageRequestHeader requestHeader = this.convertToPopMessageRequestHeader(ctx, request);
SelectableMessageQueue messageQueue = this.readQueueSelector.select(ctx, request, requestHeader);
if (messageQueue == null) {
throw new ProxyException(Code.NOT_FOUND, "no readable topic route for topic " + requestHeader.getTopic());
}
CompletableFuture<PopResult> popResultFuture = this.readConsumer.popMessage(
messageQueue.getBrokerAddr(),
messageQueue.getBrokerName(),
@@ -182,7 +186,7 @@ public class ConsumerService extends BaseService {
}
future.whenComplete((ackResult, throwable) -> {
if (ackNoMatchedMessageHook != null) {
ackNoMatchedMessageHook.beforeResponse(ackMessageRequestHeader, ackResult, throwable);
ackNoMatchedMessageHook.beforeResponse(ctx, ackMessageRequestHeader, ackResult, throwable);
}
});
}
@@ -191,7 +195,7 @@ public class ConsumerService extends BaseService {
CompletableFuture<AckMessageResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (ackMessageHook != null) {
ackMessageHook.beforeResponse(request, response, throwable);
ackMessageHook.beforeResponse(ctx, request, response, throwable);
}
});
try {
@@ -237,7 +241,7 @@ public class ConsumerService extends BaseService {
CompletableFuture<NackMessageResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (nackMessageHook != null) {
nackMessageHook.beforeResponse(request, response, throwable);
nackMessageHook.beforeResponse(ctx, request, response, throwable);
}
});
try {
@@ -258,7 +262,6 @@ public class ConsumerService extends BaseService {
}
})
.exceptionally(throwable -> {
throwable.printStackTrace();
future.completeExceptionally(throwable);
return null;
});
@@ -72,14 +72,14 @@ public class ForwardClientService extends BaseService {
this.producerManager.setProducerOfflineListener(connectorManager.getTransactionHeartbeatRegisterService()::onProducerGroupOffline);
}
public void heartbeat(Context ctx, HeartbeatRequest request, ChannelManager channelManager) {
String language = InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.LANGUAGE);
public void heartbeat(Context ctx, HeartbeatRequest request) {
String language = InterceptorConstants.METADATA.get(ctx).get(InterceptorConstants.LANGUAGE);
LanguageCode languageCode = LanguageCode.valueOf(language);
String clientId = request.getClientId();
if (request.hasProducerData()) {
String producerGroup = GrpcConverter.wrapResourceWithNamespace(request.getProducerData().getGroup());
GrpcClientChannel channel = GrpcClientChannel.create(channelManager, producerGroup, clientId, pollCommandResponseManager);
GrpcClientChannel channel = GrpcClientChannel.create(ctx, channelManager, producerGroup, clientId, pollCommandResponseManager);
ClientChannelInfo clientChannelInfo = new ClientChannelInfo(channel, clientId, languageCode, MQVersion.Version.V5_0_0.ordinal());
producerManager.registerProducer(producerGroup, clientChannelInfo);
}
@@ -87,7 +87,7 @@ public class ForwardClientService extends BaseService {
if (request.hasConsumerData()) {
ConsumerData consumerData = request.getConsumerData();
String consumerGroup = GrpcConverter.wrapResourceWithNamespace(consumerData.getGroup());
GrpcClientChannel channel = GrpcClientChannel.create(channelManager, consumerGroup, clientId, pollCommandResponseManager);
GrpcClientChannel channel = GrpcClientChannel.create(ctx, channelManager, consumerGroup, clientId, pollCommandResponseManager);
ClientChannelInfo clientChannelInfo = new ClientChannelInfo(channel, clientId, languageCode, MQVersion.Version.V5_0_0.ordinal());
consumerManager.registerConsumer(
@@ -102,7 +102,7 @@ public class ForwardClientService extends BaseService {
}
}
public void unregister(Context ctx, NotifyClientTerminationRequest request, ChannelManager channelManager) {
public void unregister(Context ctx, NotifyClientTerminationRequest request) {
String clientId = request.getClientId();
if (request.hasProducerGroup()) {
@@ -164,4 +164,12 @@ public class ForwardClientService extends BaseService {
LOGGER.error("error occurred when scan not active client channels.", e);
}
}
public ConsumerManager getConsumerManager() {
return consumerManager;
}
public ProducerManager getProducerManager() {
return producerManager;
}
}
@@ -68,7 +68,7 @@ public class ProducerService extends BaseService {
CompletableFuture<SendMessageResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (sendMessageHook != null) {
sendMessageHook.beforeResponse(request, response, throwable);
sendMessageHook.beforeResponse(ctx, request, response, throwable);
}
});
@@ -140,7 +140,7 @@ public class ProducerService extends BaseService {
CompletableFuture<ForwardMessageToDeadLetterQueueResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (forwardMessageToDLQHook != null) {
forwardMessageToDLQHook.beforeResponse(request, response, throwable);
forwardMessageToDLQHook.beforeResponse(ctx, request, response, throwable);
}
});
try {
@@ -61,7 +61,7 @@ public class PullMessageService extends BaseService {
CompletableFuture<QueryOffsetResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (queryOffsetHook != null) {
queryOffsetHook.beforeResponse(request, response, throwable);
queryOffsetHook.beforeResponse(ctx, request, response, throwable);
}
});
try {
@@ -100,7 +100,7 @@ public class PullMessageService extends BaseService {
CompletableFuture<PullMessageResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (pullMessageHook != null) {
pullMessageHook.beforeResponse(request, response, throwable);
pullMessageHook.beforeResponse(ctx, request, response, throwable);
}
});
@@ -135,10 +135,11 @@ public class PullMessageService extends BaseService {
// check filterExpression is correct or not
GrpcConverter.buildSubscriptionData(GrpcConverter.wrapResourceWithNamespace(request.getPartition().getTopic()), request.getFilterExpression());
long pollTime = ctx.getDeadline()
.timeRemaining(TimeUnit.MILLISECONDS) - ConfigurationManager.getProxyConfig().getLongPollingReserveTimeInMillis();
long timeRemaining = ctx.getDeadline()
.timeRemaining(TimeUnit.MILLISECONDS);
long pollTime = timeRemaining - ConfigurationManager.getProxyConfig().getLongPollingReserveTimeInMillis();
if (pollTime <= 0) {
throw new ProxyException(Code.DEADLINE_EXCEEDED, "request has been canceled due to timeout");
pollTime = timeRemaining;
}
return GrpcConverter.buildPullMessageRequestHeader(request, pollTime);
}
@@ -97,7 +97,7 @@ public class RouteService extends BaseService {
CompletableFuture<QueryRouteResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (queryRouteHook != null) {
queryRouteHook.beforeResponse(request, response, throwable);
queryRouteHook.beforeResponse(ctx, request, response, throwable);
}
});
@@ -209,7 +209,7 @@ public class RouteService extends BaseService {
CompletableFuture<QueryAssignmentResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (queryAssignmentHook != null) {
queryAssignmentHook.beforeResponse(request, response, throwable);
queryAssignmentHook.beforeResponse(ctx, request, response, throwable);
}
});
@@ -26,11 +26,13 @@ import io.grpc.Context;
import java.util.List;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ThreadLocalRandom;
import org.apache.commons.collections.CollectionUtils;
import org.apache.rocketmq.common.protocol.header.EndTransactionRequestHeader;
import org.apache.rocketmq.proxy.channel.ChannelManager;
import org.apache.rocketmq.proxy.common.utils.ProxyUtils;
import org.apache.rocketmq.proxy.connector.ConnectorManager;
import org.apache.rocketmq.proxy.connector.ForwardProducer;
import org.apache.rocketmq.proxy.connector.transaction.TransactionId;
import org.apache.rocketmq.proxy.connector.transaction.TransactionStateCheckRequest;
import org.apache.rocketmq.proxy.connector.transaction.TransactionStateChecker;
import org.apache.rocketmq.proxy.grpc.adapter.channel.GrpcClientChannel;
@@ -54,9 +56,12 @@ public class TransactionService extends BaseService implements TransactionStateC
@Override
public void checkTransactionState(TransactionStateCheckRequest checkData) {
Context ctx = Context.current();
try {
List<String> clientIdList = this.channelManager.getClientIdList(checkData.getGroupId());
// if clientIdList's size is 0, here will throw: java.lang.IllegalArgumentException: bound must be positive
if (CollectionUtils.isEmpty(clientIdList)) {
return;
}
String clientId = clientIdList.get(ThreadLocalRandom.current().nextInt(clientIdList.size()));
GrpcClientChannel channel = GrpcClientChannel.getChannel(this.channelManager, checkData.getGroupId(), clientId);
@@ -72,11 +77,11 @@ public class TransactionService extends BaseService implements TransactionStateC
).build();
channel.writeAndFlush(response);
if (this.checkTransactionStateHook != null) {
this.checkTransactionStateHook.beforeResponse(checkData, response, null);
this.checkTransactionStateHook.beforeResponse(ctx, checkData, response, null);
}
} catch (Throwable t) {
if (this.checkTransactionStateHook != null) {
this.checkTransactionStateHook.beforeResponse(checkData, null, t);
this.checkTransactionStateHook.beforeResponse(ctx, checkData, null, t);
}
}
}
@@ -85,13 +90,16 @@ public class TransactionService extends BaseService implements TransactionStateC
CompletableFuture<EndTransactionResponse> future = new CompletableFuture<>();
future.whenComplete((response, throwable) -> {
if (endTransactionHook != null) {
endTransactionHook.beforeResponse(request, response, throwable);
endTransactionHook.beforeResponse(ctx, request, response, throwable);
}
});
try {
TransactionId handle = TransactionId.decode(request.getTransactionId());
EndTransactionRequestHeader requestHeader = this.toEndTransactionRequestHeader(ctx, request);
this.forwardProducer.endTransaction(requestHeader, ProxyUtils.DEFAULT_MQ_CLIENT_TIMEOUT);
this.forwardProducer.endTransaction(
RemotingHelper.parseSocketAddressAddr(handle.getBrokerAddr()),
requestHeader, ProxyUtils.DEFAULT_MQ_CLIENT_TIMEOUT);
future.complete(EndTransactionResponse.newBuilder()
.setCommon(ResponseBuilder.buildCommon(Code.OK, Code.OK.name()))
.build());
@@ -32,7 +32,7 @@ public class InitConfigAndLoggerTest {
public static String mockProxyHome = "/mock/rmq/proxy/home";
@Before
public void before() throws Exception {
public void before() throws Throwable {
URL mockProxyHomeURL = getClass().getClassLoader().getResource("rmq-proxy-home");
if (mockProxyHomeURL != null) {
mockProxyHome = mockProxyHomeURL.toURI().getPath();
@@ -107,7 +107,7 @@ public class LocalGrpcServiceTest extends InitConfigAndLoggerTest {
private Metadata metadata;
@Before
public void setUp() throws Exception {
public void setUp() throws Throwable {
super.before();
Mockito.when(brokerControllerMock.getSendMessageProcessor()).thenReturn(sendMessageProcessorMock);
Mockito.when(brokerControllerMock.getPopMessageProcessor()).thenReturn(popMessageProcessorMock);
@@ -16,12 +16,23 @@
*/
package org.apache.rocketmq.proxy.grpc.service.cluster;
import java.net.SocketAddress;
import java.nio.charset.StandardCharsets;
import java.util.concurrent.ThreadLocalRandom;
import java.util.concurrent.TimeUnit;
import org.apache.rocketmq.common.consumer.ReceiptHandle;
import org.apache.rocketmq.common.message.MessageAccessor;
import org.apache.rocketmq.common.message.MessageConst;
import org.apache.rocketmq.common.message.MessageExt;
import org.apache.rocketmq.proxy.config.InitConfigAndLoggerTest;
import org.apache.rocketmq.proxy.connector.ConnectorManager;
import org.apache.rocketmq.proxy.connector.DefaultForwardClient;
import org.apache.rocketmq.proxy.connector.ForwardProducer;
import org.apache.rocketmq.proxy.connector.ForwardReadConsumer;
import org.apache.rocketmq.proxy.connector.route.TopicRouteCache;
import org.apache.rocketmq.proxy.connector.ForwardWriteConsumer;
import org.apache.rocketmq.proxy.connector.transaction.TransactionHeartbeatRegisterService;
import org.apache.rocketmq.remoting.common.RemotingUtil;
import org.junit.Before;
import org.junit.Ignore;
import org.junit.runner.RunWith;
@@ -32,10 +43,10 @@ import static org.mockito.Mockito.when;
@Ignore
@RunWith(MockitoJUnitRunner.Silent.class)
public abstract class BaseServiceTest {
public abstract class BaseServiceTest extends InitConfigAndLoggerTest {
@Mock
protected ConnectorManager clientManager;
protected ConnectorManager connectorManager;
@Mock
protected DefaultForwardClient defaultClient;
@Mock
@@ -46,17 +57,53 @@ public abstract class BaseServiceTest {
protected ForwardWriteConsumer writeConsumerClient;
@Mock
protected TopicRouteCache topicRouteCache;
@Mock
protected TransactionHeartbeatRegisterService transactionHeartbeatRegisterService;
@Before
public void before() throws Throwable {
when(clientManager.getDefaultForwardClient()).thenReturn(defaultClient);
when(clientManager.getForwardProducer()).thenReturn(producerClient);
when(clientManager.getForwardReadConsumer()).thenReturn(readConsumerClient);
when(clientManager.getForwardWriteConsumer()).thenReturn(writeConsumerClient);
when(clientManager.getTopicRouteCache()).thenReturn(topicRouteCache);
super.before();
when(connectorManager.getDefaultForwardClient()).thenReturn(defaultClient);
when(connectorManager.getForwardProducer()).thenReturn(producerClient);
when(connectorManager.getForwardReadConsumer()).thenReturn(readConsumerClient);
when(connectorManager.getForwardWriteConsumer()).thenReturn(writeConsumerClient);
when(connectorManager.getTopicRouteCache()).thenReturn(topicRouteCache);
when(connectorManager.getTransactionHeartbeatRegisterService()).thenReturn(transactionHeartbeatRegisterService);
beforeEach();
}
public abstract void beforeEach() throws Throwable;
protected static ReceiptHandle createReceiptHandle() {
return ReceiptHandle.builder()
.topicType(ReceiptHandle.NORMAL_TOPIC)
.brokerName("brokerName")
.retrieveTime(System.currentTimeMillis())
.invisibleTime(TimeUnit.SECONDS.toMillis(3))
.queueId(ThreadLocalRandom.current().nextInt(8))
.offset(ThreadLocalRandom.current().nextInt(1000))
.commitLogOffset(ThreadLocalRandom.current().nextInt(1000))
.build();
}
protected static MessageExt createMessageExt(String msgId, String tag) {
return createMessageExt(msgId, tag, createReceiptHandle().encode());
}
protected static MessageExt createMessageExt(String msgId, String tag, String handler) {
SocketAddress addr = RemotingUtil.string2SocketAddress("127.0.0.1:8080");
MessageExt msg = new MessageExt(0,
System.currentTimeMillis(),
addr,
System.currentTimeMillis(),
addr,
msgId);
msg.setTopic("topic");
msg.setBody("hello".getBytes(StandardCharsets.UTF_8));
MessageAccessor.putProperty(msg, MessageConst.PROPERTY_TAGS, tag);
MessageAccessor.putProperty(msg, MessageConst.PROPERTY_UNIQ_CLIENT_MESSAGE_ID_KEYIDX, msgId);
MessageAccessor.putProperty(msg, MessageConst.PROPERTY_POP_CK, handler);
return msg;
}
}
@@ -0,0 +1,138 @@
package org.apache.rocketmq.proxy.grpc.service.cluster;
import apache.rocketmq.v1.ConsumeMessageType;
import apache.rocketmq.v1.ConsumeModel;
import apache.rocketmq.v1.ConsumePolicy;
import apache.rocketmq.v1.ConsumerData;
import apache.rocketmq.v1.FilterExpression;
import apache.rocketmq.v1.FilterType;
import apache.rocketmq.v1.HeartbeatRequest;
import apache.rocketmq.v1.NotifyClientTerminationRequest;
import apache.rocketmq.v1.ProducerData;
import apache.rocketmq.v1.Resource;
import apache.rocketmq.v1.SubscriptionEntry;
import io.grpc.Context;
import io.grpc.Metadata;
import io.netty.channel.Channel;
import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.Executors;
import org.apache.rocketmq.broker.client.ClientChannelInfo;
import org.apache.rocketmq.broker.client.ConsumerGroupInfo;
import org.apache.rocketmq.common.consumer.ConsumeFromWhere;
import org.apache.rocketmq.common.protocol.heartbeat.ConsumeType;
import org.apache.rocketmq.common.protocol.heartbeat.MessageModel;
import org.apache.rocketmq.proxy.channel.ChannelManager;
import org.apache.rocketmq.proxy.grpc.adapter.channel.GrpcClientChannel;
import org.apache.rocketmq.proxy.grpc.common.PollCommandResponseManager;
import org.apache.rocketmq.proxy.grpc.interceptor.InterceptorConstants;
import org.apache.rocketmq.remoting.protocol.LanguageCode;
import org.junit.Test;
import static org.junit.Assert.*;
public class ClientServiceTest extends BaseServiceTest {
private ChannelManager channelManager = new ChannelManager();
private PollCommandResponseManager pollCommandResponseManager = new PollCommandResponseManager();
@Override
public void beforeEach() throws Throwable {
}
@Test
public void testProducerHeartbeat() {
ClientService clientService = new ClientService(
this.connectorManager,
Executors.newSingleThreadScheduledExecutor(),
this.channelManager,
this.pollCommandResponseManager);
Metadata metadata = new Metadata();
metadata.put(InterceptorConstants.LANGUAGE, "JAVA");
metadata.put(InterceptorConstants.REMOTE_ADDRESS, "127.0.0.1:8080");
metadata.put(InterceptorConstants.LOCAL_ADDRESS, "127.0.0.1:8081");
Context ctx = Context.current().withValue(InterceptorConstants.METADATA, metadata);
clientService.heartbeat(ctx, HeartbeatRequest.newBuilder()
.setClientId("clientId")
.setProducerData(ProducerData.newBuilder()
.setGroup(Resource.newBuilder()
.setName("producerGroup")
.build())
.build())
.build());
assertEquals(1, clientService.getProducerManager().getGroupChannelTable().size());
Channel channel = clientService.getProducerManager().findChannel("clientId");
assertNotNull(channel);
assertTrue(channel instanceof GrpcClientChannel);
clientService.unregister(ctx, NotifyClientTerminationRequest.newBuilder()
.setClientId("clientId")
.setProducerGroup(Resource.newBuilder()
.setName("producerGroup")
.build())
.build());
assertTrue(clientService.getProducerManager().getGroupChannelTable().isEmpty());
}
@Test
public void testConsumerHeartbeat() {
ClientService clientService = new ClientService(
this.connectorManager,
Executors.newSingleThreadScheduledExecutor(),
this.channelManager,
this.pollCommandResponseManager);
List<SubscriptionEntry> subscriptionEntryList = new ArrayList<>();
subscriptionEntryList.add(SubscriptionEntry.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setExpression(FilterExpression.newBuilder()
.setExpression("*")
.setType(FilterType.TAG)
.build())
.build());
Metadata metadata = new Metadata();
metadata.put(InterceptorConstants.LANGUAGE, "JAVA");
metadata.put(InterceptorConstants.REMOTE_ADDRESS, "127.0.0.1:8080");
metadata.put(InterceptorConstants.LOCAL_ADDRESS, "127.0.0.1:8081");
Context ctx = Context.current().withValue(InterceptorConstants.METADATA, metadata);
clientService.heartbeat(ctx, HeartbeatRequest.newBuilder()
.setClientId("clientId")
.setConsumerData(ConsumerData.newBuilder()
.setGroup(Resource.newBuilder()
.setName("consumerGroup")
.build())
.setConsumeType(ConsumeMessageType.PASSIVE)
.setConsumeModel(ConsumeModel.CLUSTERING)
.setConsumePolicy(ConsumePolicy.RESUME)
.addAllSubscriptions(subscriptionEntryList)
.build())
.build());
ClientChannelInfo clientChannelInfo = clientService.getConsumerManager().findChannel("consumerGroup", "clientId");
assertNotNull(clientChannelInfo);
assertEquals(LanguageCode.JAVA, clientChannelInfo.getLanguage());
assertEquals("clientId", clientChannelInfo.getClientId());
assertTrue(clientChannelInfo.getChannel() instanceof GrpcClientChannel);
ConsumerGroupInfo consumerGroupInfo = clientService.getConsumerManager().getConsumerGroupInfo("consumerGroup");
assertEquals(MessageModel.CLUSTERING, consumerGroupInfo.getMessageModel());
assertEquals(ConsumeFromWhere.CONSUME_FROM_LAST_OFFSET, consumerGroupInfo.getConsumeFromWhere());
assertEquals(ConsumeType.CONSUME_PASSIVELY, consumerGroupInfo.getConsumeType());
assertEquals("TAG", consumerGroupInfo.getSubscriptionTable().get("topic").getExpressionType());
assertEquals("*", consumerGroupInfo.getSubscriptionTable().get("topic").getSubString());
clientService.unregister(ctx, NotifyClientTerminationRequest.newBuilder()
.setClientId("clientId")
.setConsumerGroup(Resource.newBuilder()
.setName("consumerGroup")
.build())
.build());
assertNull(clientService.getConsumerManager().getConsumerGroupInfo("consumerGroup"));
}
}
@@ -0,0 +1,173 @@
package org.apache.rocketmq.proxy.grpc.service.cluster;
import apache.rocketmq.v1.AckMessageRequest;
import apache.rocketmq.v1.AckMessageResponse;
import apache.rocketmq.v1.FilterExpression;
import apache.rocketmq.v1.FilterType;
import apache.rocketmq.v1.NackMessageRequest;
import apache.rocketmq.v1.NackMessageResponse;
import apache.rocketmq.v1.Partition;
import apache.rocketmq.v1.ReceiveMessageRequest;
import apache.rocketmq.v1.ReceiveMessageResponse;
import apache.rocketmq.v1.Resource;
import com.google.rpc.Code;
import io.grpc.Context;
import java.util.List;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.Executors;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicReference;
import org.apache.rocketmq.client.consumer.AckResult;
import org.apache.rocketmq.client.consumer.AckStatus;
import org.apache.rocketmq.client.consumer.PopResult;
import org.apache.rocketmq.client.consumer.PopStatus;
import org.apache.rocketmq.common.consumer.ReceiptHandle;
import org.apache.rocketmq.common.message.MessageExt;
import org.apache.rocketmq.common.message.MessageQueue;
import org.apache.rocketmq.common.protocol.ResponseCode;
import org.apache.rocketmq.common.protocol.header.ChangeInvisibleTimeRequestHeader;
import org.apache.rocketmq.common.protocol.header.ConsumerSendMsgBackRequestHeader;
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.Mock;
import static org.junit.Assert.assertEquals;
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.when;
public class ConsumerServiceTest extends BaseServiceTest {
@Mock
private ReadQueueSelector readQueueSelector;
private ConsumerService consumerService;
@Override
public void beforeEach() throws Throwable {
consumerService = new ConsumerService(this.connectorManager);
consumerService.setReadQueueSelector(readQueueSelector);
}
@Test
public void testReceiveMessage() throws Exception {
SelectableMessageQueue selectableMessageQueue = new SelectableMessageQueue(
new MessageQueue("namespace%topic", "brokerName", 0), "brokerAddr");
when(readQueueSelector.select(any(), any(), any())).thenReturn(selectableMessageQueue);
List<MessageExt> messageExtList = Lists.newArrayList(
createMessageExt("msg1", "msg1"),
createMessageExt("msg2", "msg2"));
PopResult popResult = new PopResult(PopStatus.FOUND, messageExtList);
when(readConsumerClient.popMessage(anyString(), anyString(), any(), anyLong()))
.thenReturn(CompletableFuture.completedFuture(popResult));
when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr");
when(writeConsumerClient.ackMessage(anyString(), any(), anyLong()))
.thenReturn(CompletableFuture.completedFuture(new AckResult()));
Context ctx = Context.current().withDeadlineAfter(3, TimeUnit.SECONDS, Executors.newSingleThreadScheduledExecutor());
AtomicReference<String> ackHandler = new AtomicReference<>();
consumerService.setAckNoMatchedMessageHook((ctx1, request, response, t) -> ackHandler.set(request.getExtraInfo()));
ReceiveMessageResponse response = consumerService.receiveMessage(ctx,
ReceiveMessageRequest.newBuilder()
.setPartition(Partition.newBuilder()
.setTopic(Resource.newBuilder()
.setResourceNamespace("namespace")
.setName("topic")
.build())
.build())
.setFilterExpression(FilterExpression.newBuilder()
.setType(FilterType.TAG)
.setExpression("msg1")
.build())
.build()
).get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(1, response.getMessagesCount());
assertEquals("msg1", response.getMessages(0).getSystemAttribute().getMessageId());
assertEquals(ReceiptHandle.create(messageExtList.get(1)).getReceiptHandle(), ackHandler.get());
}
@Test
public void testAckMessage() throws Exception {
when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr");
AckResult ackResult = new AckResult();
ackResult.setStatus(AckStatus.OK);
when(writeConsumerClient.ackMessage(anyString(), any(), anyLong())).thenReturn(CompletableFuture.completedFuture(ackResult));
AckMessageResponse response = consumerService.ackMessage(Context.current(), AckMessageRequest.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setGroup(Resource.newBuilder()
.setName("group")
.build())
.setReceiptHandle(createReceiptHandle().encode())
.build())
.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
}
@Test
public void testNackMessageToDLQ() throws Exception {
ReceiptHandle receiptHandle = createReceiptHandle();
AtomicReference<ConsumerSendMsgBackRequestHeader> headerRef = new AtomicReference<>();
doAnswer(mock -> {
headerRef.set(mock.getArgument(1));
return CompletableFuture.completedFuture(RemotingCommand.createResponseCommand(ResponseCode.SUCCESS, ""));
}).when(producerClient).sendMessageBack(anyString(), any(), anyLong());
when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr");
NackMessageResponse response = consumerService.nackMessage(Context.current(), NackMessageRequest.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setGroup(Resource.newBuilder()
.setName("group")
.build())
.setReceiptHandle(receiptHandle.encode())
.setDeliveryAttempt(3)
.setMaxDeliveryAttempts(3)
.build())
.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(receiptHandle.getCommitLogOffset(), headerRef.get().getOffset().longValue());
}
@Test
public void testNackMessage() throws Exception {
ReceiptHandle receiptHandle = createReceiptHandle();
AtomicReference<ChangeInvisibleTimeRequestHeader> headerRef = new AtomicReference<>();
doAnswer(mock -> {
headerRef.set(mock.getArgument(2));
AckResult ackResult = new AckResult();
ackResult.setStatus(AckStatus.OK);
return CompletableFuture.completedFuture(ackResult);
}).when(writeConsumerClient).changeInvisibleTimeAsync(anyString(), anyString(), any(), anyLong());
when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr");
NackMessageResponse response = consumerService.nackMessage(Context.current(), NackMessageRequest.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setGroup(Resource.newBuilder()
.setName("group")
.build())
.setReceiptHandle(receiptHandle.encode())
.setDeliveryAttempt(1)
.setMaxDeliveryAttempts(3)
.build())
.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(receiptHandle.getOffset(), headerRef.get().getOffset().longValue());
assertEquals(receiptHandle.encode(), headerRef.get().getExtraInfo());
}
}
@@ -0,0 +1,74 @@
package org.apache.rocketmq.proxy.grpc.service.cluster;
import apache.rocketmq.v1.Broker;
import apache.rocketmq.v1.Partition;
import apache.rocketmq.v1.ReceiveMessageRequest;
import io.grpc.Context;
import org.apache.rocketmq.common.message.MessageQueue;
import org.apache.rocketmq.common.protocol.header.PopMessageRequestHeader;
import org.apache.rocketmq.proxy.connector.route.SelectableMessageQueue;
import org.junit.Test;
import static org.junit.Assert.assertNull;
import static org.junit.Assert.assertSame;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.ArgumentMatchers.isNull;
import static org.mockito.Mockito.when;
public class DefaultReadQueueSelectorTest extends BaseServiceTest {
private final String brokerAddress = "127.0.0.1:10911";
@Override
public void beforeEach() throws Throwable {
}
@Test
public void test() throws Exception {
SelectableMessageQueue messageQueue1 = new SelectableMessageQueue(
new MessageQueue("readBrokerTopicByName", "brokerName", 0), "brokerAddr1");
SelectableMessageQueue messageQueue2 = new SelectableMessageQueue(
new MessageQueue("oneReadBroker", "brokerName", 0), "brokerAddr1");
when(topicRouteCache.selectReadBrokerByName(eq("readBrokerTopicByName"), anyString())).thenReturn(messageQueue1);
when(topicRouteCache.selectOneReadBroker(eq("oneReadBroker"), isNull())).thenReturn(messageQueue2);
ReadQueueSelector readQueueSelector = new DefaultReadQueueSelector(topicRouteCache);
{
PopMessageRequestHeader requestHeader = new PopMessageRequestHeader();
requestHeader.setTopic("readBrokerTopicByName");
SelectableMessageQueue messageQueue = readQueueSelector.select(Context.current(),
ReceiveMessageRequest.newBuilder()
.setPartition(Partition.newBuilder()
.setBroker(Broker.newBuilder()
.setName("brokerName")
.build())
.build())
.build(),
requestHeader);
assertSame(messageQueue1, messageQueue);
}
{
PopMessageRequestHeader requestHeader = new PopMessageRequestHeader();
requestHeader.setTopic("oneReadBroker");
SelectableMessageQueue messageQueue = readQueueSelector.select(Context.current(),
ReceiveMessageRequest.newBuilder()
.build(),
requestHeader);
assertSame(messageQueue2, messageQueue);
}
{
PopMessageRequestHeader requestHeader = new PopMessageRequestHeader();
requestHeader.setTopic("topic");
SelectableMessageQueue messageQueue = readQueueSelector.select(Context.current(),
ReceiveMessageRequest.newBuilder()
.build(),
requestHeader);
assertNull(messageQueue);
}
}
}
@@ -21,7 +21,7 @@ import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.isNull;
import static org.mockito.Mockito.when;
public class DefaultProducerQueueSelectorTest extends BaseServiceTest {
public class DefaultWriteQueueSelectorTest extends BaseServiceTest {
@Override
public void beforeEach() throws Throwable {
@@ -33,6 +33,8 @@ import org.apache.rocketmq.common.message.MessageQueue;
import org.apache.rocketmq.proxy.connector.route.SelectableMessageQueue;
import org.apache.rocketmq.proxy.grpc.adapter.ProxyException;
import org.junit.Test;
import org.mockito.invocation.InvocationOnMock;
import org.mockito.stubbing.Answer;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertNotNull;
@@ -42,6 +44,7 @@ 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.when;
public class ProducerServiceTest extends BaseServiceTest {
@@ -71,7 +74,7 @@ public class ProducerServiceTest extends BaseServiceTest {
sendResultFuture.complete(new SendResult(SendStatus.SEND_OK, "msgId", new MessageQueue(),
1L, "txId", "offsetMsgId", "regionId"));
ProducerService producerService = new ProducerService(this.clientManager);
ProducerService producerService = new ProducerService(this.connectorManager);
producerService.setWriteQueueSelector((ctx, request, requestHeader, message) ->
new SelectableMessageQueue(new MessageQueue("namespace%topic", "brokerName", 0), "brokerAddr"));
@@ -88,7 +91,7 @@ public class ProducerServiceTest extends BaseServiceTest {
@Test
public void testSendMessageNoQueueSelect() {
ProducerService producerService = new ProducerService(this.clientManager);
ProducerService producerService = new ProducerService(this.connectorManager);
producerService.setWriteQueueSelector((ctx, request, requestHeader, message) -> null);
@@ -125,7 +128,7 @@ public class ProducerServiceTest extends BaseServiceTest {
.thenReturn(sendResultFuture);
sendResultFuture.completeExceptionally(ex);
ProducerService producerService = new ProducerService(this.clientManager);
ProducerService producerService = new ProducerService(this.connectorManager);
producerService.setWriteQueueSelector((ctx, request, requestHeader, message) ->
new SelectableMessageQueue(new MessageQueue("namespace%topic", "brokerName", 0), "brokerAddr"));
@@ -145,11 +148,11 @@ public class ProducerServiceTest extends BaseServiceTest {
public void testSendMessageWithErrorThrow() {
RuntimeException ex = new RuntimeException();
ProducerService producerService = new ProducerService(this.clientManager);
ProducerService producerService = new ProducerService(this.connectorManager);
producerService.setWriteQueueSelector((ctx, request, requestHeader, message) -> {
throw ex;
});
producerService.setSendMessageHook((request, response, t) -> {
producerService.setSendMessageHook((ctx, request, response, t) -> {
assertSame(ex, t);
});
@@ -0,0 +1,134 @@
package org.apache.rocketmq.proxy.grpc.service.cluster;
import apache.rocketmq.v1.Broker;
import apache.rocketmq.v1.FilterExpression;
import apache.rocketmq.v1.FilterType;
import apache.rocketmq.v1.Partition;
import apache.rocketmq.v1.PullMessageRequest;
import apache.rocketmq.v1.PullMessageResponse;
import apache.rocketmq.v1.QueryOffsetPolicy;
import apache.rocketmq.v1.QueryOffsetRequest;
import apache.rocketmq.v1.QueryOffsetResponse;
import apache.rocketmq.v1.Resource;
import com.google.protobuf.Timestamp;
import com.google.protobuf.util.Timestamps;
import com.google.rpc.Code;
import io.grpc.Context;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.Executors;
import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicReference;
import org.apache.rocketmq.client.consumer.PullResult;
import org.apache.rocketmq.client.consumer.PullStatus;
import org.apache.rocketmq.common.protocol.header.PullMessageRequestHeader;
import org.assertj.core.util.Lists;
import org.junit.Test;
import org.mockito.invocation.InvocationOnMock;
import org.mockito.stubbing.Answer;
import static org.junit.Assert.*;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyInt;
import static org.mockito.ArgumentMatchers.anyLong;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.Mockito.doAnswer;
import static org.mockito.Mockito.when;
public class PullMessageServiceTest extends BaseServiceTest {
private PullMessageService pullMessageService;
@Override
public void beforeEach() throws Throwable {
pullMessageService = new PullMessageService(this.connectorManager);
when(topicRouteCache.getBrokerAddr(anyString())).thenReturn("brokerAddr");
}
@Test
public void testQueryOffset() throws Exception {
Context ctx = Context.current();
when(defaultClient.getMaxOffset(anyString(), anyString(), anyInt(), anyLong())).thenReturn(CompletableFuture.completedFuture(100L));
when(defaultClient.searchOffset(anyString(), anyString(), anyInt(), anyLong(), anyLong())).thenReturn(CompletableFuture.completedFuture(50L));
QueryOffsetResponse response = pullMessageService.queryOffset(ctx, QueryOffsetRequest.newBuilder()
.setPartition(Partition.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setBroker(Broker.newBuilder().setName("brokerName").build())
.build())
.setPolicy(QueryOffsetPolicy.BEGINNING)
.build()
).get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(0, response.getOffset());
response = pullMessageService.queryOffset(ctx, QueryOffsetRequest.newBuilder()
.setPartition(Partition.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setBroker(Broker.newBuilder().setName("brokerName").build())
.build())
.setPolicy(QueryOffsetPolicy.END)
.build()
).get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(100, response.getOffset());
response = pullMessageService.queryOffset(ctx, QueryOffsetRequest.newBuilder()
.setPartition(Partition.newBuilder()
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.setBroker(Broker.newBuilder().setName("brokerName").build())
.build())
.setTimePoint(Timestamps.fromMillis(System.currentTimeMillis()))
.setPolicy(QueryOffsetPolicy.TIME_POINT)
.build()
).get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(50, response.getOffset());
}
@Test
public void testPullMessage() throws Exception {
AtomicReference<PullMessageRequestHeader> headerRef = new AtomicReference<>();
PullResult pullResult = new PullResult(
PullStatus.FOUND,
3,
0,
10,
Lists.newArrayList(
createMessageExt("msg1", "msg1"),
createMessageExt("msg2", "msg2")
)
);
doAnswer(mock -> {
headerRef.set(mock.getArgument(1));
return CompletableFuture.completedFuture(pullResult);
}).when(readConsumerClient).pullMessage(anyString(), any(), anyLong());
Context ctx = Context.current().withDeadlineAfter(3, TimeUnit.SECONDS, Executors.newSingleThreadScheduledExecutor());
PullMessageResponse response = pullMessageService.pullMessage(ctx, PullMessageRequest.newBuilder()
.setPartition(Partition.newBuilder()
.setBroker(Broker.newBuilder()
.setName("brokerName")
.build())
.setTopic(Resource.newBuilder()
.setName("topic")
.build())
.build())
.setFilterExpression(FilterExpression.newBuilder()
.setExpression("msg1")
.setType(FilterType.TAG)
.build())
.build())
.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(1, response.getMessagesCount());
assertEquals("msg1", response.getMessages(0).getSystemAttribute().getMessageId());
}
}
@@ -35,15 +35,20 @@ import java.util.ArrayList;
import java.util.HashMap;
import java.util.List;
import java.util.concurrent.CompletableFuture;
import org.apache.rocketmq.client.exception.MQClientException;
import org.apache.rocketmq.common.constant.PermName;
import org.apache.rocketmq.common.protocol.ResponseCode;
import org.apache.rocketmq.common.protocol.route.BrokerData;
import org.apache.rocketmq.common.protocol.route.QueueData;
import org.apache.rocketmq.proxy.grpc.adapter.ProxyMode;
import org.apache.rocketmq.common.protocol.route.TopicRouteData;
import org.apache.rocketmq.proxy.connector.route.MessageQueueWrapper;
import org.apache.rocketmq.proxy.grpc.common.ProxyMode;
import org.junit.Test;
import static org.assertj.core.api.Assertions.assertThat;
import static org.junit.Assert.assertEquals;
import static org.junit.Assert.assertNull;
import static org.mockito.Mockito.when;
public class RouteServiceTest extends BaseServiceTest {
private String brokerAddress = "127.0.0.1:10911";
@@ -57,7 +62,9 @@ public class RouteServiceTest extends BaseServiceTest {
.build();
@Override
public void beforeEach() {
public void beforeEach() throws Exception {
TopicRouteData routeData = new TopicRouteData();
List<BrokerData> brokerDataList = new ArrayList<>();
BrokerData brokerData = new BrokerData();
brokerData.setCluster("cluster");
@@ -75,13 +82,21 @@ public class RouteServiceTest extends BaseServiceTest {
queueData.setReadQueueNums(8);
queueData.setBrokerName("brokerName");
queueDataList.add(queueData);
routeData.setBrokerDatas(brokerDataList);
routeData.setQueueDatas(queueDataList);
MessageQueueWrapper messageQueueWrapper = new MessageQueueWrapper("topic", routeData);
when(this.topicRouteCache.getMessageQueue("topic")).thenReturn(messageQueueWrapper);
when(this.topicRouteCache.getMessageQueue("notExistTopic")).thenThrow(new MQClientException(ResponseCode.TOPIC_NOT_EXIST, ""));
}
@Test
public void testGenPartitionFromQueueData() throws Exception {
// test queueData with 8 read queues, 8 write queues, and rw permission, expect 8 rw queues.
QueueData queueDataWith8R8WPermRW = mockQueueData(8, 8, PermName.PERM_READ | PermName.PERM_WRITE);
List<Partition> partitionWith8R8WPermRW = RouteService.genPartitionFromQueueData(queueDataWith8R8WPermRW, MOCK_TOPIC, MOCK_BROKER);
List<Partition> partitionWith8R8WPermRW = RouteService.genPartitionFromQueueData(queueDataWith8R8WPermRW, MOCK_TOPIC, MOCK_BROKER);
assertThat(partitionWith8R8WPermRW.size()).isEqualTo(8);
assertThat(partitionWith8R8WPermRW.stream().filter(a -> a.getPermission() == Permission.READ_WRITE).count()).isEqualTo(8);
assertThat(partitionWith8R8WPermRW.stream().filter(a -> a.getPermission() == Permission.READ).count()).isEqualTo(0);
@@ -89,7 +104,7 @@ public class RouteServiceTest extends BaseServiceTest {
// test queueData with 8 read queues, 8 write queues, and read only permission, expect 8 read only queues.
QueueData queueDataWith8R8WPermR = mockQueueData(8, 8, PermName.PERM_READ);
List<Partition> partitionWith8R8WPermR = RouteService.genPartitionFromQueueData(queueDataWith8R8WPermR, MOCK_TOPIC, MOCK_BROKER);
List<Partition> partitionWith8R8WPermR = RouteService.genPartitionFromQueueData(queueDataWith8R8WPermR, MOCK_TOPIC, MOCK_BROKER);
assertThat(partitionWith8R8WPermR.size()).isEqualTo(8);
assertThat(partitionWith8R8WPermR.stream().filter(a -> a.getPermission() == Permission.READ).count()).isEqualTo(8);
assertThat(partitionWith8R8WPermR.stream().filter(a -> a.getPermission() == Permission.READ_WRITE).count()).isEqualTo(0);
@@ -97,7 +112,7 @@ public class RouteServiceTest extends BaseServiceTest {
// test queueData with 8 read queues, 8 write queues, and write only permission, expect 8 write only queues.
QueueData queueDataWith8R8WPermW = mockQueueData(8, 8, PermName.PERM_WRITE);
List<Partition> partitionWith8R8WPermW = RouteService.genPartitionFromQueueData(queueDataWith8R8WPermW, MOCK_TOPIC, MOCK_BROKER);
List<Partition> partitionWith8R8WPermW = RouteService.genPartitionFromQueueData(queueDataWith8R8WPermW, MOCK_TOPIC, MOCK_BROKER);
assertThat(partitionWith8R8WPermW.size()).isEqualTo(8);
assertThat(partitionWith8R8WPermW.stream().filter(a -> a.getPermission() == Permission.WRITE).count()).isEqualTo(8);
assertThat(partitionWith8R8WPermW.stream().filter(a -> a.getPermission() == Permission.READ_WRITE).count()).isEqualTo(0);
@@ -105,7 +120,7 @@ public class RouteServiceTest extends BaseServiceTest {
// test queueData with 8 read queues, 0 write queues, and rw permission, expect 8 read only queues.
QueueData queueDataWith8R0WPermRW = mockQueueData(8, 0, PermName.PERM_READ | PermName.PERM_WRITE);
List<Partition> partitionWith8R0WPermRW = RouteService.genPartitionFromQueueData(queueDataWith8R0WPermRW, MOCK_TOPIC, MOCK_BROKER);
List<Partition> partitionWith8R0WPermRW = RouteService.genPartitionFromQueueData(queueDataWith8R0WPermRW, MOCK_TOPIC, MOCK_BROKER);
assertThat(partitionWith8R0WPermRW.size()).isEqualTo(8);
assertThat(partitionWith8R0WPermRW.stream().filter(a -> a.getPermission() == Permission.READ).count()).isEqualTo(8);
assertThat(partitionWith8R0WPermRW.stream().filter(a -> a.getPermission() == Permission.READ_WRITE).count()).isEqualTo(0);
@@ -113,7 +128,7 @@ public class RouteServiceTest extends BaseServiceTest {
// test queueData with 4 read queues, 8 write queues, and rw permission, expect 4 rw queues and 4 write only queues.
QueueData queueDataWith4R8WPermRW = mockQueueData(4, 8, PermName.PERM_READ | PermName.PERM_WRITE);
List<Partition> partitionWith4R8WPermRW = RouteService.genPartitionFromQueueData(queueDataWith4R8WPermRW, MOCK_TOPIC, MOCK_BROKER);
List<Partition> partitionWith4R8WPermRW = RouteService.genPartitionFromQueueData(queueDataWith4R8WPermRW, MOCK_TOPIC, MOCK_BROKER);
assertThat(partitionWith4R8WPermRW.size()).isEqualTo(8);
assertThat(partitionWith4R8WPermRW.stream().filter(a -> a.getPermission() == Permission.WRITE).count()).isEqualTo(4);
assertThat(partitionWith4R8WPermRW.stream().filter(a -> a.getPermission() == Permission.READ_WRITE).count()).isEqualTo(4);
@@ -131,8 +146,8 @@ public class RouteServiceTest extends BaseServiceTest {
}
@Test
public void testLocalModeQueryRoute() {
RouteService routeService = new RouteService(ProxyMode.LOCAL, this.clientManager);
public void testLocalModeQueryRoute() throws Exception {
RouteService routeService = new RouteService(ProxyMode.LOCAL, this.connectorManager);
CompletableFuture<QueryRouteResponse> future = routeService.queryRoute(Context.current(), QueryRouteRequest.newBuilder()
.setEndpoints(Endpoints.newBuilder()
.addAddresses(Address.newBuilder()
@@ -145,20 +160,16 @@ public class RouteServiceTest extends BaseServiceTest {
.setName("topic")
.build())
.build());
try {
QueryRouteResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(8, response.getPartitionsCount());
assertEquals(HostAndPort.fromString(brokerAddress).getHost(), response.getPartitions(0).getBroker()
.getEndpoints().getAddresses(0).getHost());
} catch (Exception e) {
assertNull(e);
}
QueryRouteResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(8, response.getPartitionsCount());
assertEquals(HostAndPort.fromString(brokerAddress).getHost(), response.getPartitions(0).getBroker()
.getEndpoints().getAddresses(0).getHost());
}
@Test
public void testQueryRouteWithInvalidEndpoints() {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.clientManager);
public void testQueryRouteWithInvalidEndpoints() throws Exception {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.connectorManager);
CompletableFuture<QueryRouteResponse> future = routeService.queryRoute(Context.current(), QueryRouteRequest.newBuilder()
.setTopic(Resource.newBuilder()
@@ -166,17 +177,13 @@ public class RouteServiceTest extends BaseServiceTest {
.build())
.build());
try {
QueryRouteResponse response = future.get();
assertEquals(Code.INVALID_ARGUMENT.getNumber(), response.getCommon().getStatus().getCode());
} catch (Exception e) {
assertNull(e);
}
QueryRouteResponse response = future.get();
assertEquals(Code.INVALID_ARGUMENT.getNumber(), response.getCommon().getStatus().getCode());
}
@Test
public void testQueryRoute() {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.clientManager);
public void testQueryRoute() throws Exception {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.connectorManager);
CompletableFuture<QueryRouteResponse> future = routeService.queryRoute(Context.current(), QueryRouteRequest.newBuilder()
.setEndpoints(Endpoints.newBuilder()
@@ -191,20 +198,16 @@ public class RouteServiceTest extends BaseServiceTest {
.build())
.build());
try {
QueryRouteResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(8, response.getPartitionsCount());
assertEquals("host", response.getPartitions(0).getBroker()
.getEndpoints().getAddresses(0).getHost());
} catch (Exception e) {
assertNull(e);
}
QueryRouteResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(8, response.getPartitionsCount());
assertEquals("host", response.getPartitions(0).getBroker()
.getEndpoints().getAddresses(0).getHost());
}
@Test
public void testQueryRouteWhenTopicNotExist() {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.clientManager);
public void testQueryRouteWhenTopicNotExist() throws Exception {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.connectorManager);
CompletableFuture<QueryRouteResponse> future = routeService.queryRoute(Context.current(), QueryRouteRequest.newBuilder()
.setEndpoints(Endpoints.newBuilder()
@@ -219,17 +222,13 @@ public class RouteServiceTest extends BaseServiceTest {
.build())
.build());
try {
QueryRouteResponse response = future.get();
assertEquals(Code.NOT_FOUND.getNumber(), response.getCommon().getStatus().getCode());
} catch (Exception e) {
assertNull(e);
}
QueryRouteResponse response = future.get();
assertEquals(Code.NOT_FOUND.getNumber(), response.getCommon().getStatus().getCode());
}
@Test
public void testQueryAssignmentInvalidEndpoints() {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.clientManager);
public void testQueryAssignmentInvalidEndpoints() throws Exception {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.connectorManager);
CompletableFuture<QueryAssignmentResponse> future = routeService.queryAssignment(Context.current(), QueryAssignmentRequest.newBuilder()
.setTopic(
@@ -239,17 +238,13 @@ public class RouteServiceTest extends BaseServiceTest {
)
.build());
try {
QueryAssignmentResponse response = future.get();
assertEquals(Code.INVALID_ARGUMENT.getNumber(), response.getCommon().getStatus().getCode());
} catch (Exception e) {
assertNull(e);
}
QueryAssignmentResponse response = future.get();
assertEquals(Code.INVALID_ARGUMENT.getNumber(), response.getCommon().getStatus().getCode());
}
@Test
public void testLocalModeQueryAssignment() {
RouteService routeService = new RouteService(ProxyMode.LOCAL, this.clientManager);
public void testLocalModeQueryAssignment() throws Exception {
RouteService routeService = new RouteService(ProxyMode.LOCAL, this.connectorManager);
CompletableFuture<QueryAssignmentResponse> future = routeService.queryAssignment(Context.current(), QueryAssignmentRequest.newBuilder()
.setEndpoints(Endpoints.newBuilder()
@@ -268,20 +263,16 @@ public class RouteServiceTest extends BaseServiceTest {
.setClientId("clientId")
.build());
try {
QueryAssignmentResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(1, response.getAssignmentsCount());
assertEquals("brokerName", response.getAssignments(0).getPartition().getBroker().getName());
assertEquals(HostAndPort.fromString(brokerAddress).getHost(), response.getAssignments(0).getPartition().getBroker().getEndpoints().getAddresses(0).getHost());
} catch (Exception e) {
assertNull(e);
}
QueryAssignmentResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(1, response.getAssignmentsCount());
assertEquals("brokerName", response.getAssignments(0).getPartition().getBroker().getName());
assertEquals(HostAndPort.fromString(brokerAddress).getHost(), response.getAssignments(0).getPartition().getBroker().getEndpoints().getAddresses(0).getHost());
}
@Test
public void testQueryAssignment() {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.clientManager);
public void testQueryAssignment() throws Exception {
RouteService routeService = new RouteService(ProxyMode.CLUSTER, this.connectorManager);
CompletableFuture<QueryAssignmentResponse> future = routeService.queryAssignment(Context.current(), QueryAssignmentRequest.newBuilder()
.setEndpoints(Endpoints.newBuilder()
@@ -300,15 +291,11 @@ public class RouteServiceTest extends BaseServiceTest {
.setClientId("clientId")
.build());
try {
QueryAssignmentResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(1, response.getAssignmentsCount());
assertEquals("brokerName", response.getAssignments(0).getPartition().getBroker().getName());
assertEquals("host", response.getAssignments(0).getPartition().getBroker().getEndpoints().getAddresses(0).getHost());
} catch (Exception e) {
assertNull(e);
}
QueryAssignmentResponse response = future.get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(1, response.getAssignmentsCount());
assertEquals("brokerName", response.getAssignments(0).getPartition().getBroker().getName());
assertEquals("host", response.getAssignments(0).getPartition().getBroker().getEndpoints().getAddresses(0).getHost());
}
}
@@ -0,0 +1,94 @@
package org.apache.rocketmq.proxy.grpc.service.cluster;
import apache.rocketmq.v1.EndTransactionRequest;
import apache.rocketmq.v1.EndTransactionResponse;
import apache.rocketmq.v1.PollCommandResponse;
import apache.rocketmq.v1.Resource;
import com.google.rpc.Code;
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;
import org.apache.rocketmq.proxy.connector.transaction.TransactionStateCheckRequest;
import org.apache.rocketmq.proxy.grpc.adapter.channel.GrpcClientChannel;
import org.apache.rocketmq.remoting.common.RemotingHelper;
import org.assertj.core.util.Lists;
import org.junit.Test;
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.mock;
import static org.mockito.Mockito.when;
public class TransactionServiceTest extends BaseServiceTest {
private TransactionService transactionService;
@Mock
private ChannelManager channelManager;
@Override
public void beforeEach() throws Throwable {
transactionService = new TransactionService(this.connectorManager, this.channelManager);
}
@Test
public void testCheckTransactionState() {
GrpcClientChannel channel = mock(GrpcClientChannel.class);
AtomicReference<Object> 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());
TransactionId transactionId = TransactionId.genFromBrokerTransactionId(
RemotingHelper.string2SocketAddress("127.0.0.1:8080"),
"71F99B78B6E261357FA259CCA6456118", 1234, 5678);
transactionService.checkTransactionState(new TransactionStateCheckRequest(
"group",
1L,
2L,
"msgId",
transactionId,
createMessageExt("msgId", "msgId")
));
assertTrue(writeDataRef.get() instanceof PollCommandResponse);
PollCommandResponse response = (PollCommandResponse) writeDataRef.get();
assertEquals(transactionId.getProxyTransactionId(), response.getRecoverOrphanedTransactionCommand().getTransactionId());
}
@Test
public void testEndTransaction() throws Exception {
AtomicReference<EndTransactionRequestHeader> headerRef = new AtomicReference<>();
AtomicReference<String> brokerAddrRef = new AtomicReference<>();
TransactionId transactionId = TransactionId.genFromBrokerTransactionId(
RemotingHelper.string2SocketAddress("127.0.0.1:8080"),
"71F99B78B6E261357FA259CCA6456118", 1234, 5678);
doAnswer(mock -> {
brokerAddrRef.set(mock.getArgument(0));
headerRef.set(mock.getArgument(1));
return null;
}).when(producerClient).endTransaction(anyString(), any(), anyLong());
EndTransactionResponse response = transactionService.endTransaction(Context.current(), EndTransactionRequest.newBuilder()
.setGroup(Resource.newBuilder()
.setName("group")
.build())
.setTransactionId(transactionId.getProxyTransactionId())
.build()
).get();
assertEquals(Code.OK.getNumber(), response.getCommon().getStatus().getCode());
assertEquals(transactionId.getBrokerTransactionId(), headerRef.get().getTransactionId());
assertEquals("127.0.0.1:8080", brokerAddrRef.get());
}
}
@@ -49,6 +49,8 @@ import io.netty.handler.ssl.util.SelfSignedCertificate;
import java.io.IOException;
import java.security.cert.CertificateException;
import java.util.concurrent.TimeUnit;
import java.util.function.Function;
import java.util.function.Supplier;
import org.apache.rocketmq.proxy.config.ConfigurationManager;
import org.apache.rocketmq.proxy.grpc.interceptor.ContextInterceptor;
import org.apache.rocketmq.proxy.grpc.interceptor.HeaderInterceptor;
@@ -124,6 +126,24 @@ public class GrpcBaseTest extends BaseConf {
.build();
}
public SendMessageRequest buildSendDelayMessageRequest(String topic, String messageId, int delayLevel) {
// Message message;
// message.getSystemAttribute().getTimedDeliveryCase();
return SendMessageRequest.newBuilder()
.setMessage(Message.newBuilder()
.setTopic(Resource.newBuilder()
.setName(topic)
.build())
.setSystemAttribute(SystemAttribute.newBuilder()
.setMessageId(messageId)
.setPartitionId(0)
.setDelayLevel(delayLevel)
.build())
.setBody(ByteString.copyFromUtf8("123"))
.build())
.build();
}
public ReceiveMessageRequest buildReceiveMessageRequest(String group, String topic) {
return ReceiveMessageRequest.newBuilder()
.setGroup(Resource.newBuilder()
@@ -1,12 +1,17 @@
package org.apache.rocketmq.test.proxy;
import apache.rocketmq.v1.AckMessageResponse;
import apache.rocketmq.v1.Address;
import apache.rocketmq.v1.AddressScheme;
import apache.rocketmq.v1.Endpoints;
import apache.rocketmq.v1.MessagingServiceGrpc;
import apache.rocketmq.v1.QueryRouteResponse;
import apache.rocketmq.v1.ReceiveMessageResponse;
import apache.rocketmq.v1.SendMessageResponse;
import com.google.common.base.Stopwatch;
import io.grpc.Channel;
import java.net.URL;
import java.util.concurrent.TimeUnit;
import org.apache.rocketmq.proxy.config.ConfigurationManager;
import org.apache.rocketmq.proxy.grpc.GrpcMessagingProcessor;
import org.apache.rocketmq.proxy.grpc.service.ClusterGrpcService;
@@ -16,7 +21,9 @@ import org.junit.After;
import org.junit.Before;
import org.junit.Test;
import static org.apache.rocketmq.common.message.MessageClientIDSetter.createUniqID;
import static org.apache.rocketmq.proxy.config.ConfigurationManager.RMQ_PROXY_HOME;
import static org.junit.Assert.assertTrue;
public class ClusterGrpcTest extends GrpcBaseTest {
@@ -61,4 +68,40 @@ public class ClusterGrpcTest extends GrpcBaseTest {
.build()));
assertQueryRoute(response, brokerControllerList.size());
}
@Test
public void testSendReceiveMessage() {
String group = "group";
String messageId = createUniqID();
SendMessageResponse sendResponse = blockingStub.sendMessage(buildSendMessageRequest(broker1Name, messageId));
assertSendMessage(sendResponse, messageId);
ReceiveMessageResponse receiveResponse = blockingStub.withDeadlineAfter(3, TimeUnit.SECONDS)
.receiveMessage(buildReceiveMessageRequest(group, broker1Name));
assertReceiveMessage(receiveResponse, messageId);
String receiptHandle = receiveResponse.getMessages(0).getSystemAttribute().getReceiptHandle();
AckMessageResponse ackMessageResponse = blockingStub.ackMessage(buildAckMessageRequest(group, broker1Name, receiptHandle));
assertAck(ackMessageResponse);
}
@Test
public void testSendReceiveDelayMessage() {
String group = "group";
String messageId = createUniqID();
SendMessageResponse sendResponse = blockingStub.sendMessage(buildSendDelayMessageRequest(broker1Name, messageId, 2));
assertSendMessage(sendResponse, messageId);
Stopwatch stopwatch = Stopwatch.createStarted();
ReceiveMessageResponse receiveResponse = blockingStub.withDeadlineAfter(10, TimeUnit.SECONDS)
.receiveMessage(buildReceiveMessageRequest(group, broker1Name));
long rcvTime = stopwatch.elapsed(TimeUnit.SECONDS);
assertTrue(Math.abs(rcvTime - 5) < 2);
assertReceiveMessage(receiveResponse, messageId);
String receiptHandle = receiveResponse.getMessages(0).getSystemAttribute().getReceiptHandle();
AckMessageResponse ackMessageResponse = blockingStub.ackMessage(buildAckMessageRequest(group, broker1Name, receiptHandle));
assertAck(ackMessageResponse);
}
}