AutoRegister基本原理分析
- 注册信息配置
- Jar文件扫描&Class文件扫描
- 查找要代码注入的类(codeInsertToClassName)
- 查找实现指定接口和类的子类(scanInterface)
- 代码注入ASM
- 运行看下效果
注册信息配置
1、注册信息配置

image.png
2、注册信息解析
void convertConfig() {
registerInfo.each { map ->
RegisterInfo info = new RegisterInfo()
info.interfaceName = map.get('scanInterface')
def superClasses = map.get('scanSuperClasses')
if (!superClasses) {
superClasses = new ArrayList<String>()
} else if (superClasses instanceof String) {
ArrayList<String> superList = new ArrayList<>()
superList.add(superClasses)
superClasses = superList
}
info.superClassNames = superClasses
info.initClassName = map.get('codeInsertToClassName') //代码注入的类
info.initMethodName = map.get('codeInsertToMethodName') //代码注入的方法(默认为static块)
info.registerMethodName = map.get('registerMethodName') //生成的代码所调用的方法
info.registerClassName = map.get('registerClassName') //注册方法所在的类
info.include = map.get('include')
info.exclude = map.get('exclude')
info.init()
if (info.validate())
list.add(info)
else {
project.logger.error('auto register config error: scanInterface, codeInsertToClassName and registerMethodName should not be null\n' + info.toString())
}
}
if (cacheEnabled) {
checkRegisterInfo()
} else {
deleteFile(AutoRegisterHelper.getRegisterInfoCacheFile(project))
deleteFile(AutoRegisterHelper.getRegisterCacheFile(project))
}
}
Jar文件扫描&Class文件扫描
// 遍历输入文件
inputs.each { TransformInput input ->
// 遍历jar
input.jarInputs.each { JarInput jarInput ->
if (jarInput.status != Status.NOTCHANGED && cacheMap) {
cacheMap.remove(jarInput.file.absolutePath)
}
scanJar(jarInput, outputProvider, scanProcessor)
}
// 遍历目录
input.directoryInputs.each { DirectoryInput directoryInput ->
long dirTime = System.currentTimeMillis();
// 获得产物的目录
File dest = outputProvider.getContentLocation(directoryInput.name, directoryInput.contentTypes, directoryInput.scopes, Format.DIRECTORY)
String root = directoryInput.file.absolutePath
if (!root.endsWith(File.separator))
root += File.separator
//遍历目录下的每个文件
directoryInput.file.eachFileRecurse { File file ->
def path = file.absolutePath.replace(root, '')
if (file.isFile()) {
def entryName = path
if (!leftSlash) {
entryName = entryName.replaceAll("\\\\", "/")
}
scanProcessor.checkInitClass(entryName, new File(dest.absolutePath + File.separator + path))
if (scanProcessor.shouldProcessClass(entryName)) {
scanProcessor.scanClass(file)
}
}
}
long scanTime = System.currentTimeMillis();
// 处理完后拷到目标文件
FileUtils.copyDirectory(directoryInput.file, dest)
println "auto-register cost time: ${System.currentTimeMillis() - dirTime}, scan time: ${scanTime - dirTime}. path=${root}"
}
}
void scanJar(JarInput jarInput, TransformOutputProvider outputProvider, CodeScanProcessor scanProcessor) {
// 获得输入文件
File src = jarInput.file
//遍历jar的字节码类文件,找到需要自动注册的类
File dest = getDestFile(jarInput, outputProvider)
long time = System.currentTimeMillis();
if (!scanProcessor.scanJar(src, dest) //直接读取了缓存,没有执行实际的扫描
//此jar文件中不需要被注入代码
//为了避免增量编译时代码注入重复,被注入代码的jar包每次都重新复制
&& !scanProcessor.isCachedJarContainsInitClass(src.absolutePath)) {
//不需要执行文件复制,直接返回
return
}
println "auto-register cost time: " + (System.currentTimeMillis() - time) + " ms to scan jar file:" + dest.absolutePath
//复制jar文件到transform目录:build/transforms/auto-register/
FileUtils.copyFile(src, dest)
}
查找要代码注入的类(codeInsertToClassName)
找到需要插入代码的类所在Jar包,存储备用 ext.fileContainsInitClass
boolean checkInitClass(String entryName, File destFile, String srcFilePath) {
if (entryName == null || !entryName.endsWith(".class"))
return
entryName = entryName.substring(0, entryName.lastIndexOf('.'))
def found = false
infoList.each { ext ->
if (ext.initClassName == entryName) {
ext.fileContainsInitClass = destFile
if (destFile.name.endsWith(".jar")) {
addToCacheMap(null, entryName, srcFilePath)
found = true
}
}
}
return found
}
查找实现指定接口和类的子类(scanInterface)
//refer hack class when object init
boolean scanClass(InputStream inputStream, String filePath) {
ClassReader cr = new ClassReader(inputStream)
ClassWriter cw = new ClassWriter(cr, 0)
ScanClassVisitor cv = new ScanClassVisitor(Opcodes.ASM5, cw, filePath)
cr.accept(cv, ClassReader.EXPAND_FRAMES)
inputStream.close()
return cv.found
}
查找实现指定接口scanInterface或指定类scanSuperClasses的子类,收集到集合 ext.classList中,备用
class ScanClassVisitor extends ClassVisitor {
private String filePath
private def found = false
ScanClassVisitor(int api, ClassVisitor cv, String filePath) {
super(api, cv)
this.filePath = filePath
}
boolean is(int access, int flag) {
return (access & flag) == flag
}
boolean isFound() {
return found
}
void visit(int version, int access, String name, String signature,
String superName, String[] interfaces) {
super.visit(version, access, name, signature, superName, interfaces)
//抽象类、接口、非public等类无法调用其无参构造方法
if (is(access, Opcodes.ACC_ABSTRACT)
|| is(access, Opcodes.ACC_INTERFACE)
|| !is(access, Opcodes.ACC_PUBLIC)
) {
return
}
infoList.each { ext ->
if (shouldProcessThisClassForRegister(ext, name)) {
if (superName != 'java/lang/Object' && !ext.superClassNames.isEmpty()) {
for (int i = 0; i < ext.superClassNames.size(); i++) {
if (ext.superClassNames.get(i) == superName) {
// println("superClassNames--------"+name)
ext.classList.add(name) //需要把对象注入到管理类 就是fileContainsInitClass
found = true
addToCacheMap(superName, name, filePath)
return
}
}
}
if (ext.interfaceName && interfaces != null) {
interfaces.each { itName ->
if (itName == ext.interfaceName) {
ext.classList.add(name)//需要把对象注入到管理类 就是fileContainsInitClass
addToCacheMap(itName, name, filePath)
found = true
}
}
}
}
}
}
}
代码注入ASM
config.list.each { ext ->
if (ext.fileContainsInitClass) {
println('')
println("insert register code to file:" + ext.fileContainsInitClass.absolutePath)
if (ext.classList.isEmpty()) {
project.logger.error("No class implements found for interface:" + ext.interfaceName)
} else {
ext.classList.each {
println(it)
}
CodeInsertProcessor.insertInitCodeTo(ext)
}
} else {
project.logger.error("The specified register class not found:" + ext.registerClassName)
}
}
static void insertInitCodeTo(RegisterInfo extension) {
if (extension != null && !extension.classList.isEmpty()) {
CodeInsertProcessor processor = new CodeInsertProcessor(extension)
File file = extension.fileContainsInitClass
if (file.getName().endsWith('.jar'))
processor.generateCodeIntoJarFile(file)
else
processor.generateCodeIntoClassFile(file)
}
}
//处理jar包中的class代码注入
private File generateCodeIntoJarFile(File jarFile) {
if (jarFile) {
def optJar = new File(jarFile.getParent(), jarFile.name + ".opt")
if (optJar.exists())
optJar.delete()
def file = new JarFile(jarFile)
Enumeration enumeration = file.entries()
JarOutputStream jarOutputStream = new JarOutputStream(new FileOutputStream(optJar))
while (enumeration.hasMoreElements()) {
JarEntry jarEntry = (JarEntry) enumeration.nextElement()
String entryName = jarEntry.getName()
ZipEntry zipEntry = new ZipEntry(entryName)
InputStream inputStream = file.getInputStream(jarEntry)
jarOutputStream.putNextEntry(zipEntry)
if (isInitClass(entryName)) {
println('generate code into:' + entryName)
def bytes = doGenerateCode(inputStream)
jarOutputStream.write(bytes)
} else {
jarOutputStream.write(IOUtils.toByteArray(inputStream))
}
inputStream.close()
jarOutputStream.closeEntry()
}
jarOutputStream.close()
file.close()
if (jarFile.exists()) {
jarFile.delete()
}
optJar.renameTo(jarFile)
}
return jarFile
}
private byte[] doGenerateCode(InputStream inputStream) {
ClassReader cr = new ClassReader(inputStream)
ClassWriter cw = new ClassWriter(cr, 0)
ClassVisitor cv = new MyClassVisitor(Opcodes.ASM5, cw)
cr.accept(cv, ClassReader.EXPAND_FRAMES)
return cw.toByteArray()
}
根据规则生成指定的ASM字节码,具体规则参考第7课
class MyMethodVisitor extends MethodVisitor {
boolean _static;
MyMethodVisitor(int api, MethodVisitor mv, boolean _static) {
super(api, mv)
this._static = _static;
}
@Override
void visitInsn(int opcode) {
if ((opcode >= Opcodes.IRETURN && opcode <= Opcodes.RETURN)) {
extension.classList.each { name ->
if (!_static) {
//加载this
mv.visitVarInsn(Opcodes.ALOAD, 0)
}
//用无参构造方法创建一个组件实例
mv.visitTypeInsn(Opcodes.NEW, name)
mv.visitInsn(Opcodes.DUP)
mv.visitMethodInsn(Opcodes.INVOKESPECIAL, name, "<init>", "()V", false)
//调用注册方法将组件实例注册到组件库中
if (_static) {
mv.visitMethodInsn(Opcodes.INVOKESTATIC
, extension.registerClassName
, extension.registerMethodName
, "(L${extension.interfaceName};)V"
, false)
} else {
mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL
, extension.registerClassName
, extension.registerMethodName
, "(L${extension.interfaceName};)V"
, false)
}
}
}
super.visitInsn(opcode)
}
@Override
void visitMaxs(int maxStack, int maxLocals) {
super.visitMaxs(maxStack + 4, maxLocals)
}
}
运行看下效果

image.png

image.png
好了,到这里,我们ASM字节码插入课程就完结了,谢谢大家。