ref[consumer]: simplify consumer configuration and support repeated consumption of the same interface by different service providers

This commit is contained in:
godotg
2024-01-13 14:07:57 +08:00
parent 331812f836
commit 82782d71ea
9 changed files with 129 additions and 143 deletions
@@ -18,6 +18,7 @@ import com.zfoo.net.consumer.balancer.AbstractConsumerLoadBalancer;
import com.zfoo.net.consumer.balancer.IConsumerLoadBalancer;
import com.zfoo.net.packet.common.Error;
import com.zfoo.net.router.Router;
import com.zfoo.net.router.SignalBridge;
import com.zfoo.net.router.answer.AsyncAnswer;
import com.zfoo.net.router.answer.SyncAnswer;
import com.zfoo.net.router.attachment.NoAnswerAttachment;
@@ -25,12 +26,11 @@ import com.zfoo.net.router.attachment.SignalAttachment;
import com.zfoo.net.router.exception.ErrorResponseException;
import com.zfoo.net.router.exception.NetTimeOutException;
import com.zfoo.net.router.exception.UnexpectedProtocolException;
import com.zfoo.net.router.SignalBridge;
import com.zfoo.net.session.Session;
import com.zfoo.net.task.TaskBus;
import com.zfoo.protocol.ProtocolManager;
import com.zfoo.protocol.collection.CollectionUtils;
import com.zfoo.protocol.registration.ProtocolModule;
import com.zfoo.protocol.exception.RunException;
import com.zfoo.protocol.util.JsonUtils;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
@@ -53,6 +53,7 @@ public class Consumer implements IConsumer {
private static final Logger logger = LoggerFactory.getLogger(Consumer.class);
// consumer|provider -> LoadBalancer
private final Map<String, IConsumerLoadBalancer> consumerLoadBalancerMap = new HashMap<>();
@Override
@@ -69,50 +70,74 @@ public class Consumer implements IConsumer {
}
public List<Session> getSessionsByModule(ProtocolModule module) {
// find all session that can process interface/packet of protocolModule
@Override
public List<Session> findProviders(Object packet) {
var protocolModule = ProtocolManager.moduleByProtocol(packet.getClass());
var list = new ArrayList<Session>();
NetContext.getSessionManager().forEachClientSession(new java.util.function.Consumer<Session>() {
@Override
public void accept(Session session) {
if (session.getConsumerAttribute() == null || session.getConsumerAttribute().getProviderConfig() == null) {
return;
}
var providerConfig = session.getConsumerAttribute().getProviderConfig();
if (providerConfig.getProviders().stream().anyMatch(it -> it.getProtocolModule().equals(module))) {
list.add(session);
}
NetContext.getSessionManager().forEachClientSession(session -> {
var consumerAttribute = session.getConsumerAttribute();
if (consumerAttribute == null) {
return;
}
var providerConfig = consumerAttribute.getProviderConfig();
if (providerConfig == null) {
return;
}
var providers = providerConfig.getProviders();
if (providers == null) {
return;
}
if (providers.stream().noneMatch(it -> it.getProtocolModule().equals(protocolModule.getName()))) {
return;
}
list.add(session);
});
if (CollectionUtils.isEmpty(list)) {
throw new RunException("[protocol:{}] has no service that provides the [module:{}]", packet.getClass().getSimpleName(), protocolModule);
}
return list;
}
public void getProvider(Object packet) {
var protocolMModule = ProtocolManager.moduleByProtocol(packet.getClass());
}
// Select a consumer loadBalancer
@Override
public IConsumerLoadBalancer loadBalancer(ProtocolModule protocolModule) {
return consumerLoadBalancerMap.get(protocolModule);
public IConsumerLoadBalancer selectLoadBalancer(List<Session> providers, Object packet) {
// select first consumer loadBalancer
// 不同的服务提供者可能会提供同一个接口,消费者可能同时消费了这些提供了同一个接口的服务提供者,取第一个消费者的loadBalancer
IConsumerLoadBalancer loadBalancer = null;
for (var providerSession : providers) {
for (var provider : providerSession.getConsumerAttribute().getProviderConfig().getProviders()) {
if (consumerLoadBalancerMap.containsKey(provider.getProvider())) {
loadBalancer = consumerLoadBalancerMap.get(provider.getProvider());
break;
}
}
if (loadBalancer != null) {
break;
}
}
if (loadBalancer == null) {
var protocolModule = ProtocolManager.moduleByProtocol(packet.getClass());
throw new RunException("[protocol:{}] can not find any loadBalancer for the [module:{}]", packet.getClass().getSimpleName(), protocolModule);
}
return loadBalancer;
}
@Override
public void send(Object packet, Object argument) {
try {
var loadBalancer = loadBalancer(ProtocolManager.moduleByProtocol(packet.getClass()));
var session = loadBalancer.loadBalancer(packet, argument);
var taskExecutorHash = TaskBus.calTaskExecutorHash(argument);
NetContext.getRouter().send(session, packet, NoAnswerAttachment.valueOf(taskExecutorHash));
} catch (Throwable t) {
logger.error("consumer unknown exception", t);
}
var providers = findProviders(packet);
var loadBalancer = selectLoadBalancer(providers, packet);
var session = loadBalancer.selectProvider(providers, packet, argument);
var taskExecutorHash = TaskBus.calTaskExecutorHash(argument);
NetContext.getRouter().send(session, packet, NoAnswerAttachment.valueOf(taskExecutorHash));
}
@Override
public <T> SyncAnswer<T> syncAsk(Object packet, Class<T> answerClass, Object argument) throws Exception {
var loadBalancer = loadBalancer(ProtocolManager.moduleByProtocol(packet.getClass()));
var session = loadBalancer.loadBalancer(packet, argument);
var providers = findProviders(packet);
var loadBalancer = selectLoadBalancer(providers, packet);
var session = loadBalancer.selectProvider(providers, packet, argument);
// 下面的代码逻辑同Router的syncAsk,如果修改的话,记得一起修改
var clientSignalAttachment = new SignalAttachment();
@@ -150,8 +175,10 @@ public class Consumer implements IConsumer {
@Override
public <T> AsyncAnswer<T> asyncAsk(Object packet, Class<T> answerClass, Object argument) {
var loadBalancer = loadBalancer(ProtocolManager.moduleByProtocol(packet.getClass()));
var session = loadBalancer.loadBalancer(packet, argument);
var providers = findProviders(packet);
var loadBalancer = selectLoadBalancer(providers, packet);
var session = loadBalancer.selectProvider(providers, packet, argument);
var asyncAnswer = NetContext.getRouter().asyncAsk(session, packet, answerClass, argument);
// load balancer之前调用
@@ -16,9 +16,11 @@ package com.zfoo.net.consumer;
import com.zfoo.net.consumer.balancer.IConsumerLoadBalancer;
import com.zfoo.net.router.answer.AsyncAnswer;
import com.zfoo.net.router.answer.SyncAnswer;
import com.zfoo.protocol.registration.ProtocolModule;
import com.zfoo.net.session.Session;
import org.springframework.lang.Nullable;
import java.util.List;
/**
* @author godotg
*/
@@ -26,7 +28,9 @@ public interface IConsumer {
void init();
IConsumerLoadBalancer loadBalancer(ProtocolModule protocolModule);
List<Session> findProviders(Object packet);
IConsumerLoadBalancer selectLoadBalancer(List<Session> providers, Object packet);
/**
* 直接发送,不需要任何返回值
@@ -13,18 +13,8 @@
package com.zfoo.net.consumer.balancer;
import com.zfoo.net.NetContext;
import com.zfoo.net.consumer.registry.RegisterVO;
import com.zfoo.net.session.Session;
import com.zfoo.protocol.ProtocolManager;
import com.zfoo.protocol.registration.ProtocolModule;
import com.zfoo.protocol.util.StringUtils;
import java.util.ArrayList;
import java.util.List;
import java.util.Objects;
import java.util.function.Consumer;
/**
* @author godotg
*/
@@ -48,35 +38,4 @@ public abstract class AbstractConsumerLoadBalancer implements IConsumerLoadBalan
return balancer;
}
public List<Session> getSessionsByModule(ProtocolModule module) {
var list = new ArrayList<Session>();
NetContext.getSessionManager().forEachClientSession(new Consumer<Session>() {
@Override
public void accept(Session session) {
if (session.getConsumerAttribute() == null || session.getConsumerAttribute().getProviderConfig() == null) {
return;
}
var providerConfig = session.getConsumerAttribute().getProviderConfig();
if (providerConfig.getProviders().stream().anyMatch(it -> it.getProtocolModule().equals(module))) {
list.add(session);
}
}
});
return list;
}
public boolean sessionHasModule(Session session, Object packet) {
var consumerAttribute = session.getConsumerAttribute();
if (Objects.isNull(consumerAttribute)) {
return false;
}
var registerVO = (RegisterVO) consumerAttribute;
if (Objects.isNull(registerVO.getProviderConfig())) {
return false;
}
var module = ProtocolManager.moduleByProtocol(packet.getClass());
return registerVO.getProviderConfig().getProviders().contains(module);
}
}
@@ -25,6 +25,7 @@ import com.zfoo.protocol.model.Pair;
import com.zfoo.protocol.registration.ProtocolModule;
import org.springframework.lang.Nullable;
import java.util.List;
import java.util.TreeMap;
import java.util.concurrent.atomic.AtomicReferenceArray;
@@ -58,17 +59,17 @@ public class ConsistentHashConsumerLoadBalancer extends AbstractConsumerLoadBala
* @return 调用的session
*/
@Override
public Session loadBalancer(Object packet, Object argument) {
public Session selectProvider(List<Session> providers, Object packet, Object argument) {
if (argument == null) {
return RandomConsumerLoadBalancer.getInstance().loadBalancer(packet, argument);
return RandomConsumerLoadBalancer.getInstance().selectProvider(providers, packet, argument);
}
updateConsistentHashMap();
updateConsistentHashMap(providers);
var module = ProtocolManager.moduleByProtocol(packet.getClass());
var fastTreeMap = consistentHashMap.get(module.getId());
if (fastTreeMap == null) {
fastTreeMap = updateModuleToConsistentHash(module);
fastTreeMap = updateModuleToConsistentHash(providers, module);
}
if (fastTreeMap == null) {
throw new RunException("ConsistentHashLoadBalancer [protocol:{}][argument:{}], no service provides the [module:{}]", packet.getClass(), argument, module);
@@ -85,7 +86,7 @@ public class ConsistentHashConsumerLoadBalancer extends AbstractConsumerLoadBala
return session;
}
private void updateConsistentHashMap() {
private void updateConsistentHashMap(List<Session> providers) {
// 如果更新时间不匹配,则更新到最新的服务提供者
var currentClientSessionChangeId = NetContext.getSessionManager().getClientSessionChangeId();
if (currentClientSessionChangeId != lastClientSessionChangeId) {
@@ -95,7 +96,7 @@ public class ConsistentHashConsumerLoadBalancer extends AbstractConsumerLoadBala
continue;
}
var module = ProtocolManager.moduleByModuleId(i);
updateModuleToConsistentHash(module);
updateModuleToConsistentHash(providers, module);
}
lastClientSessionChangeId = currentClientSessionChangeId;
}
@@ -103,8 +104,8 @@ public class ConsistentHashConsumerLoadBalancer extends AbstractConsumerLoadBala
@Nullable
private FastTreeMapIntLong updateModuleToConsistentHash(ProtocolModule module) {
var sessionStringList = getSessionsByModule(module).stream()
private FastTreeMapIntLong updateModuleToConsistentHash(List<Session> providers, ProtocolModule module) {
var sessionStringList = providers.stream()
.map(session -> new Pair<>(session.getConsumerAttribute().toString(), session.getSid()))
.sorted((a, b) -> a.getKey().compareTo(b.getKey()))
.toList();
@@ -2,7 +2,6 @@ package com.zfoo.net.consumer.balancer;
import com.zfoo.net.NetContext;
import com.zfoo.net.session.Session;
import org.apache.curator.shaded.com.google.common.collect.Lists;
import org.apache.curator.shaded.com.google.common.util.concurrent.AtomicLongMap;
import java.util.List;
@@ -31,20 +30,20 @@ public class ConsistentHashOfMemoryConsumerLoadBalancer extends ConsistentHashCo
}
@Override
public Session loadBalancer(Object packet, Object argument) {
public Session selectProvider(List<Session> providers, Object packet, Object argument) {
if (argument instanceof Long) {
long sid = uid2sidMap.get((Long) argument);
if (sid > 0L) {
Session memorySession = NetContext.getSessionManager().getClientSession(sid);
if (null != memorySession){
if (null != memorySession) {
return memorySession;
}else {
} else {
uid2sidMap.remove((Long) argument);
}
}
}
Session loadBalancer = super.loadBalancer(packet, argument);
Session loadBalancer = super.selectProvider(providers, packet, argument);
if (argument instanceof Long){
uid2sidMap.put((Long) argument, loadBalancer.getSid());
}
@@ -17,19 +17,22 @@ import com.zfoo.net.router.attachment.SignalAttachment;
import com.zfoo.net.session.Session;
import org.springframework.lang.Nullable;
import java.util.List;
/**
* @author godotg
*/
public interface IConsumerLoadBalancer {
/**
* Select a service provider that can provide interface/packet services
* 只有一致性hash会使用这个argument参数,如果在一致性hash没有传入argument默认使用随机负载均衡
*
* @param packet 请求包
* @param argument 计算参数
* @return 一个服务提供者的session
*/
Session loadBalancer(Object packet, @Nullable Object argument);
Session selectProvider(List<Session> providers, Object packet, @Nullable Object argument);
default void beforeLoadBalancer(Session session, Object packet, SignalAttachment attachment) {
}
@@ -14,10 +14,10 @@
package com.zfoo.net.consumer.balancer;
import com.zfoo.net.session.Session;
import com.zfoo.protocol.ProtocolManager;
import com.zfoo.protocol.exception.RunException;
import com.zfoo.protocol.util.RandomUtils;
import java.util.List;
/**
* 随机负载均衡器,任选服务提供者的其中之一
*
@@ -35,15 +35,8 @@ public class RandomConsumerLoadBalancer extends AbstractConsumerLoadBalancer {
}
@Override
public Session loadBalancer(Object packet, Object argument) {
var module = ProtocolManager.moduleByProtocol(packet.getClass());
var sessions = getSessionsByModule(module);
if (sessions.isEmpty()) {
throw new RunException("RandomConsumerLoadBalancer [protocol:{}][argument:{}], no service provides the [module:{}]", packet.getClass(), argument, module);
}
return RandomUtils.randomEle(sessions);
public Session selectProvider(List<Session> providers, Object packet, Object argument) {
return RandomUtils.randomEle(providers);
}
}
@@ -432,7 +432,7 @@ public class ZookeeperRegistry implements IRegistry {
return;
}
executor.execute(() -> doCheckConsumer());
executor.execute(ThreadUtils.safeRunnable(() -> doCheckConsumer()));
}
/**
@@ -450,17 +450,17 @@ public class ZookeeperRegistry implements IRegistry {
for (var providerCache : providerHashConsumerSet) {
// 先排除已经启动的consumer
// getClientSessionMap
var consumerClientList = new ArrayList<Session>();
NetContext.getSessionManager().forEachClientSession(new Consumer<Session>() {
@Override
public void accept(Session session) {
if (session.getConsumerAttribute() != null && session.getConsumerAttribute().equals(providerCache)) {
consumerClientList.add(session);
}
}
});
@Override
public void accept(Session session) {
if (session.getConsumerAttribute() != null && session.getConsumerAttribute().equals(providerCache)) {
consumerClientList.add(session);
}
}
});
// consumerClientList大于等于1,说明消费者连接成功了
if (consumerClientList.size() == 1) {
var consumer = consumerClientList.get(0);
if (SessionUtils.isActive(consumer)) {
@@ -476,41 +476,44 @@ public class ZookeeperRegistry implements IRegistry {
continue;
}
// 自己作为消费者,要创建一个TcpClient去连接服务提供者
var client = new TcpClient(HostAndPort.valueOf(providerCache.getProviderConfig().getAddress()));
var session = client.start();
try {
// 自己作为消费者,要创建一个TcpClient去连接服务提供者
var client = new TcpClient(HostAndPort.valueOf(providerCache.getProviderConfig().getAddress()));
var session = client.start();
if (session == null) {
recheckFlag = true;
continue;
}
// 自己作为消费者,使用TcpClient连接服务提供者不成功
if (Objects.isNull(session)) {
logger.error("[consumer:{}] failed to start, wait [{}] seconds to recheck consumer", providerCache, RETRY_SECONDS);
recheckFlag = true;
} else {
// 连接上了服务提供者
session.setConsumerAttribute(providerCache);
EventBus.post(ConsumerStartEvent.valueOf(providerCache, session));
try {
var localRegisterVO = NetContext.getConfigManager().getLocalConfig().toLocalRegisterVO();
var path = CONSUMER_ROOT_PATH + StringUtils.SLASH + localRegisterVO.toConsumerString();
var stat = curator.checkExists().forPath(path);
if (Objects.isNull(stat)) {
curator.create()
.withMode(CreateMode.EPHEMERAL)
.forPath(path);
} else {
curator.setData().forPath(path);
}
} catch (Exception e) {
// 因为并不关心consumer的状态,这种失败只需要记录一个错误日志就可以了
logger.error("consumer writing to Zookeeper failed", e);
}
} catch (Throwable t) {
logger.error("[consumer:{}] failed to start, wait [{}] seconds to recheck consumer", providerCache, RETRY_SECONDS, t);
recheckFlag = true;
}
}
// 将自己的消费者消息写到 /consumer 的临时节点下
var localRegisterVO = NetContext.getConfigManager().getLocalConfig().toLocalRegisterVO();
var path = CONSUMER_ROOT_PATH + StringUtils.SLASH + localRegisterVO.toConsumerString();
try {
var stat = curator.checkExists().forPath(path);
if (Objects.isNull(stat)) {
curator.create()
.withMode(CreateMode.EPHEMERAL)
.forPath(path);
} else {
curator.setData().forPath(path);
}
} catch (Exception e) {
logger.error("consumer:[{}] writing to Zookeeper failed", path, e);
recheckFlag = true;
}
if (recheckFlag) {
SchedulerBus.schedule(() -> checkConsumer(), RETRY_SECONDS, TimeUnit.SECONDS);
}
}
/**
@@ -15,8 +15,6 @@ package com.zfoo.net.handler;
import com.zfoo.event.manager.EventBus;
import com.zfoo.net.NetContext;
import com.zfoo.net.consumer.balancer.ConsistentHashConsumerLoadBalancer;
import com.zfoo.net.consumer.balancer.IConsumerLoadBalancer;
import com.zfoo.net.core.gateway.IGatewayLoadBalancer;
import com.zfoo.net.core.gateway.model.GatewaySessionInactiveEvent;
import com.zfoo.net.packet.DecodedPacketInfo;
@@ -27,7 +25,6 @@ import com.zfoo.net.router.attachment.GatewayAttachment;
import com.zfoo.net.router.attachment.SignalAttachment;
import com.zfoo.net.session.Session;
import com.zfoo.net.util.SessionUtils;
import com.zfoo.protocol.ProtocolManager;
import com.zfoo.protocol.util.JsonUtils;
import com.zfoo.protocol.util.StringUtils;
import com.zfoo.scheduler.util.TimeUtils;
@@ -117,10 +114,10 @@ public class GatewayRouteHandler extends ServerRouteHandler {
private void forwardingPacket(Object packet, Object attachment, Object argument) {
try {
// 网关统一用 moduleid uid 获取 session
var loadBalancer = NetContext.getConsumer().loadBalancer(ProtocolManager.moduleByProtocol(packet.getClass()));
Session consumerSession = loadBalancer.loadBalancer(packet, argument);
// var consumerSession = ConsistentHashConsumerLoadBalancer.getInstance().loadBalancer(packet, argument);
NetContext.getRouter().send(consumerSession, packet, attachment);
var providers = NetContext.getConsumer().findProviders(packet);
var loadBalancer = NetContext.getConsumer().selectLoadBalancer(providers, packet);
var providerSession = loadBalancer.selectProvider(providers, packet, argument);
NetContext.getRouter().send(providerSession, packet, attachment);
} catch (Exception e) {
logger.error("An exception occurred at the gateway", e);
} catch (Throwable t) {