diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/channel/ChannelManager.java b/proxy/src/main/java/org/apache/rocketmq/proxy/channel/ChannelManager.java index 61910fb248..54a54f32be 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/channel/ChannelManager.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/channel/ChannelManager.java @@ -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 clientIdChannelMap = new ConcurrentHashMap<>(); + private final ConcurrentMap/* clientId */> groupClientIdMap = new ConcurrentHashMap<>(); public SimpleChannel createChannel() { return createChannel(anonymousChannelId()); @@ -70,6 +73,10 @@ public class ChannelManager { return clazz.cast(channel); } + public void setChannel(String clientId, T channel) { + clientIdChannelMap.put(clientId, channel); + } + public T removeChannel(String clientId, Class 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; */ diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/adapter/channel/GrpcClientChannel.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/adapter/channel/GrpcClientChannel.java index de8bfa906a..d62c5f80b0 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/adapter/channel/GrpcClientChannel.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/adapter/channel/GrpcClientChannel.java @@ -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/* clientId */> GROUP_CLIENT_IDS = new ConcurrentHashMap<>(); - private final AtomicReference> pollCommandResponseFutureRef = new AtomicReference<>(); - public GrpcClientChannel() { + private GrpcClientChannel() { super(ChannelManager.createSimpleChannelDirectly()); } + public void addClientObserver(CompletableFuture 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 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); } diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/service/LocalGrpcService.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/service/LocalGrpcService.java index 9169f191ee..c5bd4658d6 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/service/LocalGrpcService.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/service/LocalGrpcService.java @@ -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 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 endTransaction(Context ctx, EndTransactionRequest request) { + @Override + public CompletableFuture endTransaction(Context ctx, EndTransactionRequest request) { return null; } @@ -295,7 +316,25 @@ public class LocalGrpcService implements GrpcForwardService { } @Override public CompletableFuture pollCommand(Context ctx, PollCommandRequest request) { - return null; + String clientId = request.getClientId(); + CompletableFuture 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 reportThreadStackTrace(Context ctx, @@ -303,7 +342,8 @@ public class LocalGrpcService implements GrpcForwardService { return null; } - @Override public CompletableFuture reportMessageConsumptionResult(Context ctx, + @Override + public CompletableFuture reportMessageConsumptionResult(Context ctx, ReportMessageConsumptionResultRequest request) { return null; }