perf[protocol]: 更加严谨的校验check protocol

This commit is contained in:
godotg
2022-10-22 17:53:47 +08:00
parent 30104a18ad
commit b2bd66a9ec
3 changed files with 45 additions and 55 deletions
@@ -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);
@@ -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;
}
@@ -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);