fix[protocol]: scan protocol class

This commit is contained in:
sun
2023-09-08 13:12:26 +08:00
parent 7ac1dab103
commit d01aa91f6d
@@ -68,11 +68,6 @@ public class ProtocolAnalysis {
// 临时变量,启动完成就会销毁,是一个基本类型序列化器
private static Map<Class<?>, ISerializer> baseSerializerMap = new HashMap<>(128);
//临时变量,存储xml配置协议id,启动完成销毁
private static Map<String, Short> protocolNameMap = new HashMap<>(MAX_PROTOCOL_NUM);
//临时变量,每个模块定义的协议
private static Map<Byte, Set<Class<?>>> moduleDefinitionClassMap = new HashMap<>(128);
static {
// 初始化基础类型序列化器
baseSerializerMap.put(boolean.class, BooleanSerializer.INSTANCE);
@@ -161,34 +156,112 @@ public class ProtocolAnalysis {
public static synchronized void analyze(XmlProtocols xmlProtocols, GenerateOperation generateOperation) {
AssertionUtils.notNull(subProtocolIdMap, "[{}]已经初始完成,请不要重复初始化", ProtocolManager.class.getSimpleName());
var protocolDefinitionMap = new HashMap<String, Boolean>();
var protocolXmlEnhanceMap = new HashMap<Class<?>, Boolean>();
var classModuleDefinitionMap = new HashMap<Class<?>, Byte>();
var moduleDefinitionClassMap = new HashMap<Byte, Set<Class<?>>>();
// 先注册类,再注册包
for (var moduleDefinition : xmlProtocols.getModules()) {
var module = new ProtocolModule(moduleDefinition.getId(), moduleDefinition.getName());
AssertionUtils.isTrue(module.getId() > 0, "[module:{}] [id:{}] 模块必须大于等于1", module.getName(), module.getId());
AssertionUtils.isNull(modules[module.getId()], "duplicate [module:{}] [id:{}] Exception!", module.getName(), module.getId());
var moduleId = moduleDefinition.getId();
var module = new ProtocolModule(moduleId, moduleDefinition.getName());
AssertionUtils.isTrue(module.getId() > 0, "[module:{}] [id:{}] 模块必须大于等于1", module.getName(), moduleId);
AssertionUtils.isNull(modules[module.getId()], "duplicate [module:{}] [id:{}] Exception!", module.getName(), moduleId);
modules[module.getId()] = module;
if (CollectionUtils.isEmpty(moduleDefinition.getProtocols())) {
continue;
}
//模块定义所有协议
Set<Class<?>> clazzSet = moduleDefinitionClassMap.computeIfAbsent(module.getId(), k -> new HashSet<>());
var protocolModuleSet = moduleDefinitionClassMap.computeIfAbsent(module.getId(), it -> new HashSet<>());
for (var protocolDefinition : moduleDefinition.getProtocols()) {
protocolDefinitionMap.put(protocolDefinition.getLocation(), protocolDefinition.isEnhance());
protocolNameMap.put(protocolDefinition.getLocation(), protocolDefinition.getId());
var id = protocolDefinition.getId();
var location = protocolDefinition.getLocation();
var enhance = protocolDefinition.isEnhance();
// Use the class path first to obtain the class name, and search the directory if it cannot be obtained
// 优先使用类路径获取类名,获取不到才去搜索目录
Class<?> clazz = null;
try {
clazz = ClassUtils.forName(location);
} catch (Exception e) {
}
// 如果定义的是类,则需要检查一下格式
if (clazz == null) {
continue;
}
var protocolId = getProtocolIdAndCheckClass(clazz);
// 没有使用Protocol注解,则使用xml定义的protocolId
if (protocolId < 0) {
if (id < 0) {
throw new RunException("[{}] Can not find protocol id, use @Protocol annotation or specify id in xml", clazz.getSimpleName());
}
protocolId = id;
} else {
if (id >= 0 && protocolId != id) {
throw new RunException("[{}] @Protocol annotation id not equal to id in xml", clazz.getSimpleName());
}
}
initProtocolClass(protocolId, clazz);
protocolModuleSet.add(clazz);
classModuleDefinitionMap.put(clazz, moduleId);
protocolXmlEnhanceMap.put(clazz, enhance);
}
}
// 再注册包路径,扫描包开始注册
for (var moduleDefinition : xmlProtocols.getModules()) {
var moduleId = moduleDefinition.getId();
var module = modules[moduleId];
if (CollectionUtils.isEmpty(moduleDefinition.getProtocols())) {
continue;
}
//模块定义所有协议
var protocolModuleSet = moduleDefinitionClassMap.computeIfAbsent(module.getId(), it -> new HashSet<>());
for (var protocolDefinition : moduleDefinition.getProtocols()) {
var id = protocolDefinition.getId();
var location = protocolDefinition.getLocation();
var enhance = protocolDefinition.isEnhance();
// Use the class path first to obtain the class name, and search the directory if it cannot be obtained
// 优先使用类路径获取类名,获取不到才去搜索目录
Class<?> clazz = null;
try {
clazz = ClassUtils.forName(location);
} catch (Exception e) {
}
// 如果定义的是类,则需要检查一下格式
if (clazz != null) {
continue;
}
var packetClazzList = scanPackageList(protocolDefinition.getLocation());
clazzSet.addAll(packetClazzList);
for (Class<?> clazz : packetClazzList) {
var previous = classModuleDefinitionMap.put(clazz, module.getId());
if (previous != null && previous != module.getId()) {
throw new RunException("[class:{}]定义到了两个不同的[module:[{}][{}]]", clazz.getName(), previous, module.getId());
} else if (previous != null && previous == module.getId()) {
//相同模块配置重复协议忽略
// 是类路径的话一定不能指定protocol id
if (CollectionUtils.isEmpty(packetClazzList)) {
throw new RunException("can not scan any protocol class in [{}]", location);
}
if (id >= 0) {
throw new RunException("When use package location, specify protocol id in xml");
}
for (Class<?> protocolClass : packetClazzList) {
// 如果location已经指定过了,则检测格式
if (classModuleDefinitionMap.containsKey(protocolClass)) {
var xmlProtocolModuleId = classModuleDefinitionMap.get(protocolClass);
if (xmlProtocolModuleId != moduleId) {
throw new RunException("[class:{}] defined two different [module:[{}][{}]]", protocolClass.getName(), xmlProtocolModuleId, module.getId());
}
continue;
}
var protocolId = getProtocolIdAndCheckClass(clazz);
initProtocolClass(protocolId, clazz);
var protocolId = getProtocolIdAndCheckClass(protocolClass);
initProtocolClass(protocolId, protocolClass);
protocolModuleSet.add(protocolClass);
classModuleDefinitionMap.put(clazz, moduleId);
protocolXmlEnhanceMap.compute(clazz, (key, value) -> Boolean.TRUE.equals(value) || enhance);
}
}
}
@@ -203,8 +276,8 @@ public class ProtocolAnalysis {
for (Class<?> clazz : packetClazzList) {
var protocolId = ProtocolManager.protocolId(clazz);
var registration = parseProtocolRegistration(clazz, module);
// 优先使用Protocol注解指定的enhance决定是否要增强协议
if (clazz.isAnnotationPresent(Protocol.class) && clazz.getAnnotation(Protocol.class).enhance()) {
// Protocol注解或者xml任意一个定义了增强协议那么就增强协议
if ((clazz.isAnnotationPresent(Protocol.class) && clazz.getAnnotation(Protocol.class).enhance()) || protocolXmlEnhanceMap.get(clazz)) {
enhanceList.add(registration);
}
// 注册协议
@@ -215,15 +288,7 @@ public class ProtocolAnalysis {
enhance(generateOperation, enhanceList);
}
public static Set<Class<?>> scanPackageList(String location) {
// Use the class path first to obtain the class name, and search the directory if it cannot be obtained
// 优先使用类路径获取类名,获取不到才去搜索目录
try {
var clazz = ClassUtils.forName(location);
return Set.of(clazz);
} catch (Exception e) {
}
public static List<Class<?>> scanPackageList(String location) {
// 获取该路径下所有类
var clazzNameSet = new HashSet<String>();
try {
@@ -239,10 +304,7 @@ public class ProtocolAnalysis {
.filter(it -> !it.isEnum())
.distinct()
.toList();
if (CollectionUtils.isEmpty(classes)) {
throw new RunException("can not scan any protocol class in [{}]", location);
}
return new HashSet<>(classes);
return new ArrayList<>(classes);
}
private static void enhance(GenerateOperation generateOperation, List<IProtocolRegistration> enhanceList) {
@@ -296,8 +358,6 @@ public class ProtocolAnalysis {
subProtocolIdMap = null;
protocolReserved = null;
baseSerializerMap = null;
protocolNameMap = null;
moduleDefinitionClassMap = null;
EnhanceUtils.clear();
@@ -579,12 +639,6 @@ public class ProtocolAnalysis {
short protocolId = -1;
if (protocolAnnotation != null) {
protocolId = protocolAnnotation.id();
} else {
// 可能通过xml的方式注册协议,xml注册协议不需要注解
Short id = protocolNameMap.get(clazz.getName());
if (id != null) {
protocolId = id;
}
}
return protocolId;