mirror of
https://github.com/tiennm99/zfoo.git
synced 2026-08-20 22:24:27 +00:00
perf[protocol]: 更加严谨的校验check protocol
This commit is contained in:
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
+2
-2
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user