diff --git a/protocol/src/main/java/com/zfoo/protocol/registration/ProtocolAnalysis.java b/protocol/src/main/java/com/zfoo/protocol/registration/ProtocolAnalysis.java index 65c6b13a..2d2df049 100644 --- a/protocol/src/main/java/com/zfoo/protocol/registration/ProtocolAnalysis.java +++ b/protocol/src/main/java/com/zfoo/protocol/registration/ProtocolAnalysis.java @@ -68,11 +68,6 @@ public class ProtocolAnalysis { // 临时变量,启动完成就会销毁,是一个基本类型序列化器 private static Map, ISerializer> baseSerializerMap = new HashMap<>(128); - //临时变量,存储xml配置协议id,启动完成销毁 - private static Map protocolNameMap = new HashMap<>(MAX_PROTOCOL_NUM); - //临时变量,每个模块定义的协议 - private static Map>> 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(); + var protocolXmlEnhanceMap = new HashMap, Boolean>(); var classModuleDefinitionMap = new HashMap, Byte>(); + var moduleDefinitionClassMap = new HashMap>>(); + + // 先注册类,再注册包 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> 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> 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> scanPackageList(String location) { // 获取该路径下所有类 var clazzNameSet = new HashSet(); 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 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;