From b2bd66a9ec5707249556ad88bf9dcf15271d70b9 Mon Sep 17 00:00:00 2001 From: godotg Date: Sat, 22 Oct 2022 17:53:47 +0800 Subject: [PATCH] =?UTF-8?q?perf[protocol]:=20=E6=9B=B4=E5=8A=A0=E4=B8=A5?= =?UTF-8?q?=E8=B0=A8=E7=9A=84=E6=A0=A1=E9=AA=8Ccheck=20protocol?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../com/zfoo/net/router/route/PacketBus.java | 3 +- .../registration/ProtocolAnalysis.java | 93 +++++++++---------- .../protobuf/GenerateProtobufUtils.java | 4 +- 3 files changed, 45 insertions(+), 55 deletions(-) diff --git a/net/src/main/java/com/zfoo/net/router/route/PacketBus.java b/net/src/main/java/com/zfoo/net/router/route/PacketBus.java index 0c10b24a..7284d6be 100644 --- a/net/src/main/java/com/zfoo/net/router/route/PacketBus.java +++ b/net/src/main/java/com/zfoo/net/router/route/PacketBus.java @@ -28,7 +28,6 @@ import com.zfoo.protocol.IPacket; import com.zfoo.protocol.ProtocolManager; import com.zfoo.protocol.collection.ArrayUtils; import com.zfoo.protocol.exception.RunException; -import com.zfoo.protocol.registration.IProtocolRegistration; import com.zfoo.protocol.registration.ProtocolAnalysis; import com.zfoo.protocol.util.AssertionUtils; import com.zfoo.protocol.util.ReflectionUtils; @@ -121,7 +120,7 @@ public abstract class PacketBus { } try { - var protocolId = ProtocolAnalysis.getProtocolIdByClass(packetClazz); + var protocolId = ProtocolAnalysis.getProtocolIdAndCheckClass(packetClazz); // 将receiver注册到IProtocolRegistration var protocolRegistration = ProtocolManager.getProtocol(protocolId); 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 e5129b7a..5561f9f1 100644 --- a/protocol/src/main/java/com/zfoo/protocol/registration/ProtocolAnalysis.java +++ b/protocol/src/main/java/com/zfoo/protocol/registration/ProtocolAnalysis.java @@ -134,7 +134,7 @@ public class ProtocolAnalysis { for (var protocolDefinition : moduleDefinition.getProtocols()) { var location = protocolDefinition.getLocation(); var clazz = Class.forName(location); - var id = getProtocolIdByClass(clazz); + var id = getProtocolIdAndCheckClass(clazz); AssertionUtils.isTrue(id >= moduleDefinition.getMinId(), "模块[{}]中的协议[{}]的协议号必须大于或者等于[{}]", moduleDefinition.getName(), clazz.getSimpleName(), moduleDefinition.getMinId()); AssertionUtils.isTrue(id < moduleDefinition.getMaxId(), "模块[{}]中的协议[{}]的协议号必须小于[{}]", moduleDefinition.getName(), clazz.getSimpleName(), moduleDefinition.getMaxId()); @@ -153,7 +153,7 @@ public class ProtocolAnalysis { for (var protocolDefinition : moduleDefinition.getProtocols()) { var location = protocolDefinition.getLocation(); var clazz = Class.forName(location); - var id = getProtocolIdByClass(clazz); + var id = getProtocolIdAndCheckClass(clazz); var registration = parseProtocolRegistration(clazz, module); if (protocolDefinition.isEnhance()) { enhanceList.add(registration); @@ -286,7 +286,7 @@ public class ProtocolAnalysis { } private static ProtocolRegistration parseProtocolRegistration(Class clazz, ProtocolModule module) { - var protocolId = getProtocolIdByClass(clazz); + var protocolId = getProtocolIdAndCheckClass(clazz); // 对象需要被序列化的属性 var fields = customFieldOrder(clazz); @@ -387,9 +387,9 @@ public class ProtocolAnalysis { return MapField.valueOf(keyRegistration, valueRegistration, type); } else { // 是一个协议引用变量 - var referenceProtocolId = getProtocolIdByClass(field.getType()); + var referenceProtocolId = getProtocolIdAndCheckClass(field.getType()); checkSubProtocol(clazz, referenceProtocolId, field.getType()); - subProtocolIdMap.computeIfAbsent(getProtocolIdByClass(clazz), it -> new HashSet<>()).add(referenceProtocolId); + subProtocolIdMap.computeIfAbsent(getProtocolIdAndCheckClass(clazz), it -> new HashSet<>()).add(referenceProtocolId); return ObjectProtocolField.valueOf(referenceProtocolId); } } @@ -425,9 +425,9 @@ public class ProtocolAnalysis { throw new RunException("不支持数组和集合联合使用[type:{}]类型", type); } else { // 是一个协议引用变量 - var referenceProtocolId = getProtocolIdByClass(clazz); + var referenceProtocolId = getProtocolIdAndCheckClass(clazz); checkSubProtocol(clazz, referenceProtocolId, clazz); - subProtocolIdMap.computeIfAbsent(getProtocolIdByClass(currentProtocolClass), it -> new HashSet<>()).add(referenceProtocolId); + subProtocolIdMap.computeIfAbsent(getProtocolIdAndCheckClass(currentProtocolClass), it -> new HashSet<>()).add(referenceProtocolId); return ObjectProtocolField.valueOf(referenceProtocolId); } } @@ -470,7 +470,7 @@ public class ProtocolAnalysis { // 协议智能语法分析,错误的协议定义将无法启动程序并给出错误警告 //----------------------------------------------------------------------- - private static void checkProtocol(Class clazz) throws IllegalAccessException, InvocationTargetException, InstantiationException { + private static void checkProtocol(Class clazz) { // 是否为一个简单的javabean ReflectionUtils.assertIsPojoClass(clazz); // 是否实现了IPacket接口 @@ -478,22 +478,7 @@ public class ProtocolAnalysis { // 不能是泛型类 AssertionUtils.isTrue(ArrayUtils.isEmpty(clazz.getTypeParameters()), "[class:{}]不能是泛型类", clazz.getCanonicalName()); - var protocolId = getProtocolIdByClass(clazz); - // 验证protocol()方法的返回是否和PROTOCOL_ID相等 - Method protocolMethod; - try { - protocolMethod = clazz.getDeclaredMethod(PROTOCOL_METHOD); - } catch (NoSuchMethodException e) { - protocolMethod = null; - } - - if (protocolMethod != null) { - // 必须要有一个空的构造器 - Constructor constructor = ReflectionUtils.publicEmptyConstructor(clazz); - IPacket packet = (IPacket) constructor.newInstance(); - var methodReturnId = (short) protocolMethod.invoke(packet); - AssertionUtils.isTrue(methodReturnId == protocolId, "[class:{}]的protocolId方法返回的值[{}]和协议号返回值[{}]不相等", clazz.getCanonicalName(), methodReturnId, protocolId); - } + var protocolId = getProtocolIdAndCheckClass(clazz); var previous = protocolClassMap.put(protocolId, clazz); if (previous != null) { @@ -501,40 +486,46 @@ public class ProtocolAnalysis { } } - public static short getProtocolIdByClass(Class clazz) { - var protocolClass = clazz.getDeclaredAnnotation(Protocol.class); - short annoProtocolId = 0; - if (protocolClass != null && protocolClass.id() != 0) { - annoProtocolId = protocolClass.id(); - } - + public static short getProtocolIdAndCheckClass(Class clazz) { Field protocolIdField = null; try { protocolIdField = clazz.getDeclaredField(PROTOCOL_ID); } catch (NoSuchFieldException e) { - if (annoProtocolId != 0) { - return annoProtocolId; + } + Method protocolMethod = null; + try { + protocolMethod = clazz.getDeclaredMethod(PROTOCOL_METHOD); + } catch (NoSuchMethodException e) { + } + + // 必须要有一个空的构造器 + Constructor constructor = ReflectionUtils.publicEmptyConstructor(clazz); + + var protocolClass = clazz.getDeclaredAnnotation(Protocol.class); + short protocolId = -1; + if (protocolClass != null && protocolClass.id() != 0) {// 注解标注的协议号 + protocolId = protocolClass.id(); + AssertionUtils.isTrue(protocolIdField == null && protocolMethod == null, "[class:{}]已经使用了注解标注协议号,不能再使用protocolId()方法和[{}]字段", clazz.getCanonicalName(), PROTOCOL_ID); + } else if (protocolIdField != null || protocolMethod != null) { // 字段标注的协议号 + AssertionUtils.isTrue(protocolIdField != null, "[class:{}]协议序列号[{}]不存在", clazz.getCanonicalName(), PROTOCOL_ID); + AssertionUtils.isTrue(Modifier.isPublic(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被public修饰", clazz.getCanonicalName(), PROTOCOL_ID); + AssertionUtils.isTrue(Modifier.isStatic(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被static修饰", clazz.getCanonicalName(), PROTOCOL_ID); + AssertionUtils.isTrue(Modifier.isFinal(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被final修饰", clazz.getCanonicalName(), PROTOCOL_ID); + AssertionUtils.isTrue(Modifier.isTransient(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被transient修饰", clazz.getCanonicalName(), PROTOCOL_ID); + AssertionUtils.isTrue(clazz.getSimpleName().matches("[a-zA-Z0-9_]*"), "[class:{}]的命名只能包含字母,数字,下划线", clazz.getCanonicalName(), PROTOCOL_ID); + + ReflectionUtils.makeAccessible(protocolIdField); + protocolId = (short) ReflectionUtils.getField(protocolIdField, null); + // 验证protocol()方法的返回是否和PROTOCOL_ID相等 + if (protocolMethod != null) { + var packet = (IPacket) ReflectionUtils.newInstance(constructor); + var methodReturnId = (short) ReflectionUtils.invokeMethod(packet, protocolMethod); + AssertionUtils.isTrue(methodReturnId == protocolId, "[class:{}]的protocolId方法返回的值[{}]和协议号返回值[{}]不相等", clazz.getCanonicalName(), methodReturnId, protocolId); } + } else { + // 可能通过xml的方式注册协议,xml注册协议不需要注解和PROTOCOL_ID协议字段号 } - // 是否被public修饰 - AssertionUtils.isTrue(protocolIdField != null, "[class:{}]协议序列号[{}]没有被public修饰", clazz.getCanonicalName(), PROTOCOL_ID); - // 是否被public修饰 - AssertionUtils.isTrue(Modifier.isPublic(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被public修饰", clazz.getCanonicalName(), PROTOCOL_ID); - // 是否被static修饰 - AssertionUtils.isTrue(Modifier.isStatic(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被static修饰", clazz.getCanonicalName(), PROTOCOL_ID); - // 是否被final修饰 - AssertionUtils.isTrue(Modifier.isFinal(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被final修饰", clazz.getCanonicalName(), PROTOCOL_ID); - // 是否被transient修饰 - AssertionUtils.isTrue(Modifier.isTransient(protocolIdField.getModifiers()), "[class:{}]协议序列号[{}]没有被transient修饰", clazz.getCanonicalName(), PROTOCOL_ID); - // 命名只能包含字母,数字,下划线 - AssertionUtils.isTrue(clazz.getSimpleName().matches("[a-zA-Z0-9_]*"), "[class:{}]的命名只能包含字母,数字,下划线", clazz.getCanonicalName(), PROTOCOL_ID); - - ReflectionUtils.makeAccessible(protocolIdField); - var protocolId = (short) ReflectionUtils.getField(protocolIdField, null); - if (annoProtocolId != 0) { - AssertionUtils.isTrue(annoProtocolId == protocolId, "[class:{}]协议序列号[{}]:[{}]与注解协议号[{}]值不相等", clazz.getCanonicalName(), PROTOCOL_ID, protocolId, annoProtocolId); - } return protocolId; } diff --git a/protocol/src/main/java/com/zfoo/protocol/serializer/protobuf/GenerateProtobufUtils.java b/protocol/src/main/java/com/zfoo/protocol/serializer/protobuf/GenerateProtobufUtils.java index e3a84454..6b8cae30 100644 --- a/protocol/src/main/java/com/zfoo/protocol/serializer/protobuf/GenerateProtobufUtils.java +++ b/protocol/src/main/java/com/zfoo/protocol/serializer/protobuf/GenerateProtobufUtils.java @@ -124,7 +124,7 @@ public abstract class GenerateProtobufUtils { for (var protos : xmlProtobuf.getProtos()) { for (var protocol : protos.getProtocols()) { var protocolClass = Class.forName(protocol.getLocation()); - var protocolId = ProtocolAnalysis.getProtocolIdByClass(protocolClass); + var protocolId = ProtocolAnalysis.getProtocolIdAndCheckClass(protocolClass); var protocolRegistration = ProtocolManager.getProtocol(protocolId); if (allGenerateProtocols.contains(protocolRegistration)) { @@ -191,7 +191,7 @@ public abstract class GenerateProtobufUtils { for (var protocol : protos.getProtocols()) { var protocolClass = Class.forName(protocol.getLocation()); - var protocolId = ProtocolAnalysis.getProtocolIdByClass(protocolClass); + var protocolId = ProtocolAnalysis.getProtocolIdAndCheckClass(protocolClass); var protocolRegistration = ProtocolManager.getProtocol(protocolId); builder.append("// id = ").append(protocolId).append(LS);