[ISSUE #3949] Implement pollCommand

This commit is contained in:
zhouxiang
2022-07-13 11:29:11 +08:00
parent 89e0727b74
commit c456ea7883
3 changed files with 63 additions and 25 deletions
@@ -20,9 +20,11 @@ package org.apache.rocketmq.proxy.channel;
import com.google.common.base.Strings;
import io.grpc.Context;
import java.util.Iterator;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ConcurrentMap;
import java.util.concurrent.CopyOnWriteArrayList;
import java.util.function.Supplier;
import org.apache.rocketmq.common.constant.LoggerName;
import org.apache.rocketmq.proxy.config.ConfigurationManager;
@@ -34,6 +36,7 @@ import org.slf4j.LoggerFactory;
public class ChannelManager {
private static final Logger LOGGER = LoggerFactory.getLogger(LoggerName.GRPC_LOGGER_NAME);
private final ConcurrentMap<String, SimpleChannel> clientIdChannelMap = new ConcurrentHashMap<>();
private final ConcurrentMap<String /* group */, List<String>/* clientId */> groupClientIdMap = new ConcurrentHashMap<>();
public SimpleChannel createChannel() {
return createChannel(anonymousChannelId());
@@ -70,6 +73,10 @@ public class ChannelManager {
return clazz.cast(channel);
}
public <T extends SimpleChannel> void setChannel(String clientId, T channel) {
clientIdChannelMap.put(clientId, channel);
}
public <T extends SimpleChannel> T removeChannel(String clientId, Class<T> clazz) {
SimpleChannel channel = clientIdChannelMap.remove(clientId);
if (channel == null) {
@@ -94,6 +101,11 @@ public class ChannelManager {
return new SimpleChannel(null, clientHost, localAddress, ConfigurationManager.getProxyConfig().getChannelExpiredInSeconds());
}
public void addGroupClientId(String group, String clientId) {
groupClientIdMap.computeIfAbsent(group, k -> new CopyOnWriteArrayList<>())
.add(clientId);
}
/**
* Scan and remove inactive mocking channels; Scan and clean expired requests;
*/
@@ -18,46 +18,32 @@ package org.apache.rocketmq.proxy.grpc.adapter.channel;
import apache.rocketmq.v1.PollCommandResponse;
import io.netty.channel.ChannelFuture;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.CopyOnWriteArrayList;
import java.util.concurrent.atomic.AtomicReference;
import org.apache.rocketmq.proxy.channel.ChannelManager;
import org.apache.rocketmq.proxy.channel.SimpleChannel;
public class GrpcClientChannel extends SimpleChannel {
private static final Map<String /* group */, List<String>/* clientId */> GROUP_CLIENT_IDS = new ConcurrentHashMap<>();
private final AtomicReference<CompletableFuture<PollCommandResponse>> pollCommandResponseFutureRef = new AtomicReference<>();
public GrpcClientChannel() {
private GrpcClientChannel() {
super(ChannelManager.createSimpleChannelDirectly());
}
public void addClientObserver(CompletableFuture<PollCommandResponse> future) {
this.pollCommandResponseFutureRef.set(future);
}
public static GrpcClientChannel create(ChannelManager channelManager, String group, String clientId) {
GrpcClientChannel channel = channelManager.createChannel(
buildKey(group, clientId),
GrpcClientChannel::new,
GrpcClientChannel.class);
GROUP_CLIENT_IDS.compute(group, (groupKey, clientIds) -> {
if (clientIds == null) {
clientIds = new CopyOnWriteArrayList<>();
}
clientIds.add(clientId);
return clientIds;
});
channelManager.addGroupClientId(group, clientId);
return channel;
}
public static void addClientObserver(ChannelManager channelManager, String group, String clientId, CompletableFuture<PollCommandResponse> future) {
GrpcClientChannel channel = getChannel(channelManager, group, clientId);
channel.pollCommandResponseFutureRef.set(future);
}
public static GrpcClientChannel getChannel(ChannelManager channelManager, String group, String clientId) {
return channelManager.getChannel(buildKey(group, clientId), GrpcClientChannel.class);
}
@@ -50,6 +50,7 @@ import apache.rocketmq.v1.ReportMessageConsumptionResultRequest;
import apache.rocketmq.v1.ReportMessageConsumptionResultResponse;
import apache.rocketmq.v1.ReportThreadStackTraceRequest;
import apache.rocketmq.v1.ReportThreadStackTraceResponse;
import apache.rocketmq.v1.Resource;
import apache.rocketmq.v1.SendMessageRequest;
import apache.rocketmq.v1.SendMessageResponse;
import com.google.rpc.Code;
@@ -75,6 +76,7 @@ import org.apache.rocketmq.proxy.channel.SimpleChannel;
import org.apache.rocketmq.proxy.channel.SimpleChannelHandlerContext;
import org.apache.rocketmq.proxy.config.ConfigurationManager;
import org.apache.rocketmq.proxy.grpc.adapter.InvocationContext;
import org.apache.rocketmq.proxy.grpc.adapter.channel.GrpcClientChannel;
import org.apache.rocketmq.proxy.grpc.adapter.channel.ReceiveMessageChannel;
import org.apache.rocketmq.proxy.grpc.adapter.channel.SendMessageChannel;
import org.apache.rocketmq.proxy.grpc.adapter.handler.ReceiveMessageResponseHandler;
@@ -111,7 +113,24 @@ public class LocalGrpcService implements GrpcForwardService {
languageCode = LanguageCode.valueOf(language);
HeartbeatData heartbeatData = Converter.buildHeartbeatData(request);
Channel channel = channelManager.createChannel();
CompletableFuture<HeartbeatResponse> future = new CompletableFuture<>();
String groupName;
switch (request.getClientDataCase()) {
case PRODUCER_DATA: {
groupName = Converter.getResourceNameWithNamespace(request.getProducerData().getGroup());
break;
}
case CONSUMER_DATA: {
groupName = Converter.getResourceNameWithNamespace(request.getConsumerData().getGroup());
break;
}
default: {
future.completeExceptionally(new IllegalArgumentException("Wrong client data type"));
return future;
}
}
GrpcClientChannel channel = GrpcClientChannel.create(channelManager, groupName, request.getClientId());
SimpleChannelHandlerContext simpleChannelHandlerContext = new SimpleChannelHandlerContext(channel);
RemotingCommand command = RemotingCommand.createRequestCommand(RequestCode.HEART_BEAT, null);
command.setLanguage(languageCode);
@@ -122,7 +141,8 @@ public class LocalGrpcService implements GrpcForwardService {
RemotingCommand response = this.brokerController.getClientManageProcessor()
.heartBeat(simpleChannelHandlerContext, command);
HeartbeatResponse heartbeatResponse = ResponseBuilder.buildHeartbeatResponse(response);
return CompletableFuture.completedFuture(heartbeatResponse);
future.complete(heartbeatResponse);
return future;
}
@Override
@@ -282,7 +302,8 @@ public class LocalGrpcService implements GrpcForwardService {
return future;
}
@Override public CompletableFuture<EndTransactionResponse> endTransaction(Context ctx, EndTransactionRequest request) {
@Override
public CompletableFuture<EndTransactionResponse> endTransaction(Context ctx, EndTransactionRequest request) {
return null;
}
@@ -295,7 +316,25 @@ public class LocalGrpcService implements GrpcForwardService {
}
@Override public CompletableFuture<PollCommandResponse> pollCommand(Context ctx, PollCommandRequest request) {
return null;
String clientId = request.getClientId();
CompletableFuture<PollCommandResponse> future = new CompletableFuture<>();
switch (request.getGroupCase()) {
case PRODUCER_GROUP:
Resource producerGroup = request.getProducerGroup();
String producerGroupName = Converter.getResourceNameWithNamespace(producerGroup);
GrpcClientChannel producerChannel = GrpcClientChannel.getChannel(channelManager, producerGroupName, clientId);
producerChannel.addClientObserver(future);
break;
case CONSUMER_GROUP:
Resource consumerGroup = request.getConsumerGroup();
String consumerGroupName = Converter.getResourceNameWithNamespace(consumerGroup);
GrpcClientChannel consumerChannel = GrpcClientChannel.getChannel(channelManager, consumerGroupName, clientId);
consumerChannel.addClientObserver(future);
break;
default:
break;
}
return future;
}
@Override public CompletableFuture<ReportThreadStackTraceResponse> reportThreadStackTrace(Context ctx,
@@ -303,7 +342,8 @@ public class LocalGrpcService implements GrpcForwardService {
return null;
}
@Override public CompletableFuture<ReportMessageConsumptionResultResponse> reportMessageConsumptionResult(Context ctx,
@Override
public CompletableFuture<ReportMessageConsumptionResultResponse> reportMessageConsumptionResult(Context ctx,
ReportMessageConsumptionResultRequest request) {
return null;
}