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 d8666fa5b5..bdbfe949dd 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 @@ -23,6 +23,7 @@ import java.util.Iterator; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; +import java.util.function.Supplier; import org.apache.rocketmq.common.constant.LoggerName; import org.apache.rocketmq.proxy.configuration.ConfigurationManager; import org.apache.rocketmq.proxy.grpc.adapter.channel.SendMessageChannel; @@ -35,21 +36,44 @@ public class ChannelManager { private final ConcurrentMap clientIdChannelMap = new ConcurrentHashMap<>(); public SimpleChannel createChannel() { - final String clientId = anonymousChannelId(); + return createChannel(anonymousChannelId()); + } + + public SimpleChannel createChannel(String clientId) { + return createChannel(clientId, ChannelManager::createSimpleChannelDirectly, SimpleChannel.class); + } + + public T createChannel(String clientId, Supplier creator, Class clazz) { if (Strings.isNullOrEmpty(clientId)) { LOGGER.warn("ClientId is unexpected null or empty"); - return createChannelInner(); + return creator.get(); } if (!clientIdChannelMap.containsKey(clientId)) { - clientIdChannelMap.putIfAbsent(clientId, createChannelInner()); + clientIdChannelMap.putIfAbsent(clientId, creator.get()); } - SimpleChannel channel = clientIdChannelMap.get(clientId); + T channel = clazz.cast(clientIdChannelMap.get(clientId)); channel.updateLastAccessTime(); return channel; } + public T getChannel(String clientId, Class clazz) { + SimpleChannel channel = clientIdChannelMap.get(clientId); + if (channel == null) { + return null; + } + return clazz.cast(channel); + } + + public T removeChannel(String clientId, Class clazz) { + SimpleChannel channel = clientIdChannelMap.remove(clientId); + if (channel == null) { + return null; + } + return clazz.cast(channel); + } + private String anonymousChannelId() { final String clientHost = InterceptorConstants.METADATA.get(Context.current()) .get(InterceptorConstants.REMOTE_ADDRESS); @@ -58,7 +82,7 @@ public class ChannelManager { return clientHost + "@" + localAddress; } - private SimpleChannel createChannelInner() { + public static SimpleChannel createSimpleChannelDirectly() { final String clientHost = InterceptorConstants.METADATA.get(Context.current()) .get(InterceptorConstants.REMOTE_ADDRESS); final String localAddress = InterceptorConstants.METADATA.get(Context.current()) 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 new file mode 100644 index 0000000000..a13d0327fe --- /dev/null +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/adapter/channel/GrpcClientChannel.java @@ -0,0 +1,82 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.rocketmq.proxy.grpc.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(SimpleChannel simpleChannel) { + super(simpleChannel); + } + + public static GrpcClientChannel create(ChannelManager channelManager, String group, String clientId) { + GrpcClientChannel channel = channelManager.createChannel( + buildKey(group, clientId), + () -> new GrpcClientChannel(ChannelManager.createSimpleChannelDirectly()), + GrpcClientChannel.class); + + GROUP_CLIENT_IDS.compute(group, (groupKey, clientIds) -> { + if (clientIds == null) { + clientIds = new CopyOnWriteArrayList<>(); + } + clientIds.add(clientId); + return clientIds; + }); + 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); + } + + public static GrpcClientChannel removeChannel(ChannelManager channelManager, String group, String clientId) { + return channelManager.removeChannel(buildKey(group, clientId), GrpcClientChannel.class); + } + + private static String buildKey(String group, String clientId) { + return group + "@" + clientId; + } + + @Override + public ChannelFuture writeAndFlush(Object msg) { + CompletableFuture future = pollCommandResponseFutureRef.get(); + if (msg instanceof PollCommandResponse) { + PollCommandResponse response = (PollCommandResponse) msg; + future.complete(response); + } + return super.writeAndFlush(msg); + } +}