← 返回

深入字节码:ASM 和 Javac Plugin 如何真正解析 Lambda

📅 2026-06-10 | 🏷️ ASM, 字节码, Javac Plugin, JVM | 📖 阅读约 30 分钟


前言

上一篇的两种方案(SerializedLambda、APT)有一个共同的限制:不能解析 Lambda 表达式,只能解析方法引用。

// ✅ 能解析 — 方法引用,有明确的方法名
Getter<User, Integer> getter = User::getAge;

// ❌ 不能解析 — Lambda 表达式,方法名是 "lambda$main$0"
Predicate<User> p = u -> u.getAge() > 18;

这一篇的两种方案可以直接解析 u -> u.getAge() > 18,但代价是需要深入 JVM 底层。


方案 3: ASM 字节码分析

原理

Lambda 表达式编译后,编译器会生成一个私有静态方法,这个方法包含了 Lambda 的实际逻辑。比如:

// 源码
Predicate<User> p = u -> u.getAge() > 18;

// 编译后,编译器在当前类中生成了一个私有方法:
// private static boolean lambda$main$0(User u) {
//     return u.getAge() > 18;
// }

这个方法的字节码是:

// lambda$main$0(User)Z
// 参数: User 对象(ALOAD_0)
// 返回: boolean

ALOAD_0                    // 加载局部变量 0(即参数 u)
INVOKEVIRTUAL User.getAge  // 调用 u.getAge(),返回 int
                           // 栈: [18](假设 getAge 返回 18)
BIPUSH 18                  // 把常量 18 压入操作数栈
                           // 栈: [18, 18]
IF_ICMPLE L1               // 比较栈顶两个 int:
                           //   如果 18 <= 18,跳转到 L1
                           //   否则继续执行
                           // 栈: []
ICONST_1                   // 压入 1(true)
                           // 栈: [1]
GOTO L2                    // 跳转到 L2
L1:
ICONST_0                   // 压入 0(false)
                           // 栈: [0]
L2:
IRETURN                    // 返回栈顶的 boolean 值

关键:我们可以通过 SerializedLambda 拿到这个方法的名字(“lambda$main$0”),然后用 ASM 读取这个方法的字节码,逐条解析指令,重建表达式树。

整体流程

Lambda 表达式: u -> u.getAge() > 18
    ↓
编译器生成: lambda$main$0(User)Z
    ↓
SerializedLambda.getImplMethodName() → "lambda$main$0"
    ↓
Class.forName(当前类) → 加载 .class 字节码
    ↓
ASM ClassReader 遍历类的所有方法
    ↓
找到 name.equals("lambda$main$0") 的方法
    ↓
自定义 MethodVisitor 逐条解析指令:
  ALOAD_0        → 标记为参数引用
  INVOKEVIRTUAL  → 识别为 getter 调用,提取属性名 "age"
  BIPUSH 18      → 识别为常量 18
  IF_ICMPLE      → 识别为比较运算,取反得到 GT
    ↓
组装: BinaryExpression(Property("age"), GT, Constant(18))

实现代码

/**
 * ASM 字节码 Lambda 解析器。
 *
 * 核心思路:
 * 1. 通过 SerializedLambda 获取 Lambda 方法的类名和方法名
 * 2. 用 ASM 读取该类的字节码
 * 3. 遍历所有方法,找到目标方法
 * 4. 用自定义的 MethodVisitor 逐条解析字节码指令
 * 5. 从指令序列中重建 Expression 树
 */
public class BytecodeLambdaParser {

    /**
     * 解析 Lambda 表达式为 Expression。
     *
     * @param predicate Lambda 表达式(必须是 Serializable 的)
     * @param entityClass 实体类类型(用于推断属性类型)
     */
    public <T> Expression parse(SerializablePredicate<T> predicate,
                                 Class<T> entityClass) {
        // Step 1: 获取 SerializedLambda
        SerializedLambda lambda = toSerializedLambda(predicate);
        String implClassName = lambda.getImplClass().replace('/', '.');
        String implMethodName = lambda.getImplMethodName();
        String implMethodSignature = lambda.getImplMethodSignature();

        // Step 2: 加载类的字节码
        // getResourceAsStream 从 classpath 中读取 .class 文件
        String resource = "/" + implClassName.replace('.', '/') + ".class";
        byte[] classBytes;
        try (InputStream is = getClass().getResourceAsStream(resource)) {
            classBytes = is.readAllBytes();
        } catch (IOException e) {
            throw new IllegalStateException("无法加载字节码: " + implClassName, e);
        }

        // Step 3: 用 ASM 解析目标方法
        // ClassReader 会遍历 .class 文件中的所有方法
        LambdaMethodVisitor visitor = new LambdaMethodVisitor(entityClass);
        ClassReader cr = new ClassReader(classBytes);
        cr.accept(new ClassVisitor(Opcodes.ASM9) {
            @Override
            public MethodVisitor visitMethod(int access, String name,
                    String descriptor, String signature, String[] exceptions) {
                // 找到目标方法(通过方法名和签名匹配)
                if (name.equals(implMethodName)
                    && descriptor.equals(implMethodSignature)) {
                    return visitor; // 返回自定义的 MethodVisitor
                }
                return null; // 其他方法跳过
            }
        }, ClassReader.SKIP_DEBUG);

        Expression result = visitor.getResult();
        if (result == null) {
            throw new IllegalStateException("无法从字节码中解析出表达式");
        }
        return result;
    }

    /**
     * 自定义 MethodVisitor — 逐条解析字节码指令,构建 Expression 树。
     *
     * 解析策略:
     * - 维护一个"虚拟栈",记录每条指令 push/pop 的值
     * - getter 调用(INVOKEVIRTUAL 无参方法)→ push PropertyExpression
     * - 常量指令(BIPUSH/SIPUSH/LDC)→ push ConstantExpression
     * - 条件跳转(IF_ICMPxx)→ 从栈中 pop 两个值,记录比较运算符
     * - visitEnd 时组装最终结果
     */
    private class LambdaMethodVisitor extends MethodVisitor {
        private final Class<?> entityClass;

        // 虚拟栈:存放 Expression 和常量值
        private final List<Object> stack = new ArrayList<>();

        // 指令日志(调试用)
        private final List<String> instructionLog = new ArrayList<>();

        // 待处理的比较运算符
        private Operator pendingComparison = null;
        private Expression comparisonLeft = null;
        private Expression comparisonRight = null;

        // 最终结果
        private Expression result;

        LambdaMethodVisitor(Class<?> entityClass) {
            super(Opcodes.ASM9);
            this.entityClass = entityClass;
        }

        Expression getResult() { return result; }

        /**
         * 处理方法调用指令。
         *
         * INVOKEVIRTUAL 指令的格式:
         *   INVOKEVIRTUAL owner.name(descriptor)
         *
         * 对于 getter 调用:
         *   INVOKEVIRTUAL User.getAge()I
         *   - owner = "User"
         *   - name = "getAge"
         *   - descriptor = "()I"(无参数,返回 int)
         */
        @Override
        public void visitMethodInsn(int opcode, String owner, String name,
                                     String descriptor, boolean isInterface) {
            if (opcode == Opcodes.INVOKEVIRTUAL
                || opcode == Opcodes.INVOKEINTERFACE) {

                // 判断是否是 getter:无参数、有返回值、不是 void
                boolean isGetter = descriptor.startsWith("()")
                    && !descriptor.equals("()V");

                // 排除自动拆箱方法(JDK 的 Integer.intValue() 等)
                // JDK 26 的字节码中,Integer 类型的 getter 返回后会自动拆箱
                boolean isUnboxing = ("intValue".equals(name)
                    || "longValue".equals(name)
                    || "doubleValue".equals(name)
                    || "booleanValue".equals(name))
                    && owner.startsWith("java/lang/");

                if (isGetter && !isUnboxing) {
                    // 这是一个 getter 调用
                    // 从方法名推断属性名
                    String propName = resolvePropertyName(name);
                    Class<?> propType = resolveAsmType(
                        descriptor.substring(2)); // 去掉 "()"

                    stack.add(new PropertyExpression(propName, propType));
                    logInstruction("INVOKEVIRTUAL " + owner + "." + name
                        + " → Property(" + propName + ")");
                    return;
                }

                // equals 方法调用 → EQ 比较
                // u.getName().equals("Tom") 编译后:
                //   ALOAD_0
                //   INVOKEVIRTUAL User.getName()String   ← getter
                //   LDC "Tom"                            ← 常量
                //   INVOKEVIRTUAL String.equals(Object)   ← equals 调用
                if ("equals".equals(name) && stack.size() >= 2) {
                    // 栈顶是 equals 的参数("Tom")
                    Object arg = stack.remove(stack.size() - 1);
                    // 下一个是 getter 的结果(Property("name"))
                    Object target = stack.remove(stack.size() - 1);

                    if (arg instanceof ConstantExpression ce
                        && target instanceof PropertyExpression pe) {
                        stack.add(new BinaryExpression(pe, Operator.EQ, ce));
                        logInstruction("INVOKEVIRTUAL equals → EQ");
                    }
                    return;
                }
            }

            // 忽略其他静态方法调用(如 Integer.valueOf 装箱操作)
            if (opcode == Opcodes.INVOKESTATIC) {
                logInstruction("INVOKESTATIC " + owner + "." + name
                    + " (ignored)");
            }
        }

        /**
         * 处理整数常量指令。
         * BIPUSH <byte> — 把一个 byte 范围的整数压入栈
         * SIPUSH <short> — 把一个 short 范围的整数压入栈
         */
        @Override
        public void visitIntInsn(int opcode, int operand) {
            stack.add(new ConstantExpression(operand, int.class));
            logInstruction((opcode == Opcodes.BIPUSH ? "BIPUSH " : "SIPUSH ")
                + operand);
        }

        /**
         * 处理 LDC 指令。
         * LDC <constant> — 从常量池中加载常量(String, int, long, float, double)
         */
        @Override
        public void visitLdcInsn(Object value) {
            Class<?> type = value.getClass();
            if (value instanceof Integer) type = int.class;
            else if (value instanceof Long) type = long.class;
            else if (value instanceof Float) type = float.class;
            else if (value instanceof Double) type = double.class;

            stack.add(new ConstantExpression(value, type));
            logInstruction("LDC " + value);
        }

        /**
         * 处理条件跳转指令。
         *
         * IF_ICMPxx 指令的语义是:如果条件成立,则跳转。
         * 但我们要的是"不跳转"时的语义(即取反)。
         *
         * 例如 IF_ICMPLE(如果 <= 则跳转):
         * - 跳转时:age <= 18 → 走 false 分支
         * - 不跳转时:age > 18 → 走 true 分支
         * - 所以 IF_ICMPLE 对应的比较运算符是 GT(取反)
         */
        @Override
        public void visitJumpInsn(int opcode, Label label) {
            Operator op = switch (opcode) {
                case Opcodes.IF_ICMPEQ -> Operator.NE;   // == 取反为 !=
                case Opcodes.IF_ICMPNE -> Operator.EQ;   // != 取反为 ==
                case Opcodes.IF_ICMPLT -> Operator.GE;   // <  取反为 >=
                case Opcodes.IF_ICMPGE -> Operator.LT;   // >= 取反为 <
                case Opcodes.IF_ICMPGT -> Operator.LE;   // >  取反为 <=
                case Opcodes.IF_ICMPLE -> Operator.GT;   // <= 取反为 >
                default -> null;
            };

            if (op != null && stack.size() >= 2) {
                // IF_ICMPxx 会从栈中 pop 两个值
                Object right = stack.remove(stack.size() - 1);
                Object left = stack.remove(stack.size() - 1);

                if (left instanceof PropertyExpression pe
                    && right instanceof ConstantExpression ce) {
                    pendingComparison = op;
                    comparisonLeft = pe;
                    comparisonRight = ce;
                    logInstruction("IF_ICMPxx → " + op);
                }
            }
        }

        /**
         * 方法结束时,组装最终结果。
         */
        @Override
        public void visitEnd() {
            logInstruction("END");

            if (pendingComparison != null && comparisonRight != null) {
                // 有比较运算 → 组装 BinaryExpression
                result = new BinaryExpression(
                    comparisonLeft, pendingComparison, comparisonRight);
            } else if (!stack.isEmpty()) {
                // 栈顶是 Expression(如 equals 的结果)
                Object top = stack.get(stack.size() - 1);
                if (top instanceof Expression expr) {
                    result = expr;
                }
            }
        }

        private Class<?> resolveAsmType(String desc) {
            return switch (desc) {
                case "I" -> int.class;
                case "J" -> long.class;
                case "F" -> float.class;
                case "D" -> double.class;
                case "Z" -> boolean.class;
                case "Ljava/lang/String;" -> String.class;
                default -> Object.class;
            };
        }

        private String resolvePropertyName(String methodName) {
            if (methodName.startsWith("get") && methodName.length() > 3) {
                String rest = methodName.substring(3);
                return Character.toLowerCase(rest.charAt(0))
                    + rest.substring(1);
            }
            if (methodName.startsWith("is") && methodName.length() > 2) {
                String rest = methodName.substring(2);
                return Character.toLowerCase(rest.charAt(0))
                    + rest.substring(1);
            }
            return methodName;
        }

        private void logInstruction(String instruction) {
            instructionLog.add(instruction);
        }
    }
}

风险

ASM 路线有两个主要风险:

  1. 依赖 JVM 字节码结构 — 不同 JVM 实现(HotSpot、OpenJ9、GraalVM)生成的字节码可能不同
  2. JDK 升级可能改变编译器行为 — 比如 JDK 26 引入的自动拆箱优化就让我的解析器第一次跑挂了

方案 4: Javac Plugin

原理

Java 编译器(javac)在编译源码时,会先经过几个阶段:

源码 (.java)
  ↓ javac parser(词法分析 + 语法分析)
AST(抽象语法树)
  ↓ javac analyzer(语义分析、类型检查)
带类型信息的 AST
  ↓ javac code gen(字节码生成)
字节码 (.class)

Javac Plugin 可以在 AST 阶段拦截,直接访问源码的树形结构。这和 C# 的 Expression Tree 本质相同,都是在编译期访问源码的 AST。

AST 节点类型

编译器的 AST 包含以下节点类型,每种对应源码中的一种语法结构:

节点类型对应的源码示例
LambdaExpressionTreeLambda 表达式u -> u.getAge() > 18
BinaryTree二元运算>, <, ==, &&, `
MethodInvocationTree方法调用getAge(), equals("Tom")
LiteralTree字面量常量18, "Tom", true
MemberSelectTree成员选择u.getAge(选中 User 的 getAge)
IdentifierTree标识符u(Lambda 参数引用)
ParenthesizedTree括号表达式(a > b)
UnaryTree一元运算!deleted(逻辑非)

对于 u -> u.getAge() > 18,AST 结构是:

LambdaExpressionTree
├── parameters: [u]
└── body: BinaryTree (GREATER_THAN)
     ├── left: MethodInvocationTree
     │   ├── methodSelect: MemberSelectTree
     │   │   ├── expression: IdentifierTree (u)
     │   │   └── identifier: "getAge"
     │   └── arguments: []
     └── right: LiteralTree (18)

实现代码

/**
 * 用 javac 编译器 API 解析 Java 源码中的 Lambda AST。
 *
 * 原理:
 * 1. 把源码字符串交给 javac 编译器
 * 2. 调用 task.parse() 获取 AST
 * 3. 用 TreeScanner 遍历 AST,找到 LambdaExpressionTree
 * 4. 递归解析 Lambda 的 body,构建 Expression 树
 */
public class LambdaAstParser {

    /**
     * 从 Java 源码中解析 Lambda 表达式。
     *
     * @param javaSource 包含 Lambda 的 Java 源码
     * @return Expression 树
     */
    public static Expression parseLambdaFromSource(String javaSource) {
        // Step 1: 创建 javac 编译器
        JavaCompiler compiler = JavacTool.create();
        DiagnosticCollector<JavaFileObject> diagnostics =
            new DiagnosticCollector<>();

        // Step 2: 创建内存中的源码文件对象
        // javac 通常从文件系统读取 .java 文件
        // 这里我们用一个内存中的字符串模拟文件
        StandardJavaFileManager fileManager =
            compiler.getStandardFileManager(diagnostics, null, null);
        JavaFileObject sourceFile = new SimpleJavaFileObject(
            URI.create("string:///Demo.java"),
            JavaFileObject.Kind.SOURCE
        ) {
            @Override
            public CharSequence getCharContent(boolean ignoreEncodingErrors) {
                return javaSource;
            }
        };

        // Step 3: 创建编译任务
        JavacTask task = (JavacTask) compiler.getTask(
            null,       // 输出流
            fileManager,
            diagnostics,
            null,       // 编译选项
            null,       // 注解处理器
            List.of(sourceFile)
        );

        try {
            // Step 4: 解析 AST
            // task.parse() 返回编译器生成的 AST 树
            Iterable<? extends Tree> trees = task.parse();

            // Step 5: 遍历 AST,找到 Lambda 表达式
            LambdaFinder finder = new LambdaFinder();
            for (Tree tree : trees) {
                finder.scan(tree, null);
            }

            if (finder.lastLambda != null) {
                // Step 6: 解析 Lambda body
                return parseLambdaTree(finder.lastLambda);
            }

            throw new IllegalStateException("源码中未找到 Lambda 表达式");
        } catch (Exception e) {
            throw new IllegalStateException("解析源码失败: " + e.getMessage(), e);
        }
    }

    /**
     * AST 遍历器 — 查找 LambdaExpressionTree 节点。
     *
     * SimpleTreeVisitor 是 javac 提供的 AST 遍历工具,
     * 我们只需要重写 visitLambda 方法来捕获 Lambda 节点。
     */
    private static class LambdaFinder extends SimpleTreeVisitor<Void, Void> {
        LambdaExpressionTree lastLambda;

        @Override
        public Void visitLambda(LambdaExpressionTree node, Void unused) {
            lastLambda = node; // 记录找到的 Lambda
            return super.visitLambda(node, unused); // 继续遍历子节点
        }
    }

    /**
     * 递归解析 Lambda 的 body。
     *
     * Lambda 的 body 可能是一个表达式(单行 Lambda),
     * 也可能是一个代码块(多行 Lambda)。
     * 这里只处理单行 Lambda 的情况。
     */
    private static Expression parseLambdaTree(LambdaExpressionTree lambda) {
        return parseExpression(lambda.getBody());
    }

    /**
     * 递归解析 AST 表达式节点。
     *
     * 这是整个 Javac Plugin 方案的核心方法。
     * 根据 AST 节点的类型,递归地构建 Expression 树。
     */
    private static Expression parseExpression(Tree tree) {
        Tree.Kind kind = tree.getKind();

        // === 二元比较运算 ===
        // age > 18 → BinaryTree(GREATER_THAN)
        if (kind == Tree.Kind.GREATER_THAN
            || kind == Tree.Kind.LESS_THAN
            || kind == Tree.Kind.EQUAL_TO
            || kind == Tree.Kind.NOT_EQUAL_TO
            || kind == Tree.Kind.GREATER_THAN_EQUAL
            || kind == Tree.Kind.LESS_THAN_EQUAL) {

            BinaryTree binary = (BinaryTree) tree;
            Operator op = switch (kind) {
                case GREATER_THAN -> Operator.GT;
                case LESS_THAN -> Operator.LT;
                case EQUAL_TO -> Operator.EQ;
                case NOT_EQUAL_TO -> Operator.NE;
                case GREATER_THAN_EQUAL -> Operator.GE;
                case LESS_THAN_EQUAL -> Operator.LE;
                default -> throw new IllegalArgumentException(
                    "不支持的比较: " + kind);
            };

            // 递归解析左右操作数
            Expression left = parseExpression(binary.getLeftOperand());
            Expression right = parseExpression(binary.getRightOperand());
            return new BinaryExpression(left, op, right);
        }

        // === 逻辑 AND ===
        if (kind == Tree.Kind.AND) {
            BinaryTree binary = (BinaryTree) tree;
            Expression left = parseExpression(binary.getLeftOperand());
            Expression right = parseExpression(binary.getRightOperand());
            return new LogicalExpression(LogicalOperator.AND,
                List.of(left, right));
        }

        // === 逻辑 OR ===
        if (kind == Tree.Kind.OR) {
            BinaryTree binary = (BinaryTree) tree;
            Expression left = parseExpression(binary.getLeftOperand());
            Expression right = parseExpression(binary.getRightOperand());
            return new LogicalExpression(LogicalOperator.OR,
                List.of(left, right));
        }

        // === 逻辑 NOT ===
        if (kind == Tree.Kind.LOGICAL_COMPLEMENT) {
            UnaryTree unary = (UnaryTree) tree;
            Expression operand = parseExpression(unary.getExpression());
            return new NotExpression(operand);
        }

        // === 字面量常量 ===
        // 18 → LiteralTree(value=18)
        // "Tom" → LiteralTree(value="Tom")
        if (tree instanceof LiteralTree literal) {
            Object value = literal.getValue();
            Class<?> type = value.getClass();
            if (value instanceof Integer) type = int.class;
            else if (value instanceof Long) type = long.class;
            else if (value instanceof Boolean) type = boolean.class;
            return new ConstantExpression(value, type);
        }

        // === 方法调用 ===
        // u.getAge() → MethodInvocationTree
        // name.equals("Tom") → MethodInvocationTree
        if (tree instanceof MethodInvocationTree method) {
            if (method.getMethodSelect() instanceof MemberSelectTree select) {
                String methodName = select.getIdentifier().toString();

                // equals 方法 → EQ 比较
                if ("equals".equals(methodName)
                    && method.getArguments().size() == 1) {
                    Expression target = parseExpression(select.getExpression());
                    Expression arg = parseExpression(
                        method.getArguments().get(0));
                    if (target instanceof PropertyExpression pe
                        && arg instanceof ConstantExpression ce) {
                        return new BinaryExpression(pe, Operator.EQ, ce);
                    }
                }

                // getter 调用(无参数、以 get/is 开头)
                if (method.getArguments().isEmpty() && isGetter(methodName)) {
                    String propName = resolvePropertyName(methodName);
                    return new PropertyExpression(propName, Object.class);
                }
            }
        }

        // === 成员选择 ===
        // u.getAge → MemberSelectTree(expression=u, identifier="getAge")
        if (tree instanceof MemberSelectTree select) {
            String name = select.getIdentifier().toString();
            if (isGetter(name)) {
                return new PropertyExpression(
                    resolvePropertyName(name), Object.class);
            }
            return new PropertyExpression(name, Object.class);
        }

        // === 标识符 ===
        // u(Lambda 参数引用)→ IdentifierTree(name="u")
        if (tree instanceof IdentifierTree id) {
            return new PropertyExpression(id.getName().toString(),
                Object.class);
        }

        // === 括号表达式 ===
        if (tree instanceof ParenthesizedTree paren) {
            return parseExpression(paren.getExpression());
        }

        throw new IllegalArgumentException(
            "不支持的 AST 节点: " + kind
            + " (" + tree.getClass().getSimpleName() + ")");
    }

    private static boolean isGetter(String name) {
        return (name.startsWith("get") && name.length() > 3)
            || (name.startsWith("is") && name.length() > 2);
    }

    private static String resolvePropertyName(String methodName) {
        if (methodName.startsWith("get") && methodName.length() > 3) {
            String rest = methodName.substring(3);
            return Character.toLowerCase(rest.charAt(0))
                + rest.substring(1);
        }
        if (methodName.startsWith("is") && methodName.length() > 2) {
            String rest = methodName.substring(2);
            return Character.toLowerCase(rest.charAt(0))
                + rest.substring(1);
        }
        return methodName;
    }
}

为什么是更接近 C# 的方案

C# 的 Expression Tree 和 Java 的 Javac Plugin 都利用编译期语法结构,但产物和可用阶段不同。

区别在于:

  • C# 是语言级特性,编译器原生支持
  • Java 需要通过 Plugin 机制扩展,而且只能在编译期使用

两条路线对比

特性ASM 字节码Javac Plugin
信息来源.class 字节码.java 源码 AST
信息完整度只有指令序列完整 AST + 类型信息
Lambda 支持
升级风险高(字节码结构可能变)低(AST API 更稳定)
运行时动态✅(可以解析运行时 Lambda)❌(只在编译期有效)
实现复杂度★★★★★★★
JDK 配置无特殊配置需要 –add-exports

下一篇:JDK Proxy — 不解析 Lambda,但工程上更实用的方案,天然支持链式 DSL。