Java实现字节码增强案例

wen java案例 2

本文目录导读:

Java实现字节码增强案例

  1. 基础环境准备
  2. 演示用的目标类
  3. Javassist实现字节码增强
  4. ASM实现字节码增强
  5. Byte Buddy实现字节码增强
  6. 完整测试示例
  7. 实际应用案例:AOP日志框架
  8. 使用注意事项

我来为你详细介绍Java字节码增强的实现案例,字节码增强是在Java类加载时或运行时修改字节码的技术,常见应用包括AOP、性能监控、热部署等。

基础环境准备

<!-- pom.xml 依赖 -->
<dependencies>
    <!-- Javassist -->
    <dependency>
        <groupId>org.javassist</groupId>
        <artifactId>javassist</artifactId>
        <version>3.29.2-GA</version>
    </dependency>
    <!-- ASM -->
    <dependency>
        <groupId>org.ow2.asm</groupId>
        <artifactId>asm</artifactId>
        <version>9.6</version>
    </dependency>
    <!-- Byte Buddy -->
    <dependency>
        <groupId>net.bytebuddy</groupId>
        <artifactId>byte-buddy</artifactId>
        <version>1.14.9</version>
    </dependency>
</dependencies>

演示用的目标类

// 目标类:用户服务
public class UserService {
    public String getUserInfo(String userId) {
        // 模拟业务逻辑
        delay(100);
        return "User-" + userId;
    }
    public void updateUser(String userId, String name) {
        // 模拟业务逻辑
        delay(200);
        System.out.println("Update user: " + userId + ", name: " + name);
    }
    private void delay(long millis) {
        try {
            Thread.sleep(millis);
        } catch (InterruptedException e) {
            Thread.currentThread().interrupt();
        }
    }
}

Javassist实现字节码增强

import javassist.*;
import java.lang.reflect.Method;
/**
 * Javassist字节码增强示例
 * 功能:方法耗时统计、参数日志记录
 */
public class JavassistEnhancer {
    // 静态代理方式
    public static Class<?> enhanceUserService() throws Exception {
        ClassPool pool = ClassPool.getDefault();
        // 获取目标类
        CtClass ctClass = pool.get("UserService");
        // 添加一个方法耗时统计的注解(可选)
        ctClass.addAnnotation(new javassist.bytecode.annotation.Annotation(
            "java.lang.Deprecated", ctClass.getClassFile().getConstPool()
        ));
        // 增强方法:添加耗时统计
        CtMethod[] methods = ctClass.getDeclaredMethods();
        for (CtMethod method : methods) {
            if (method.isEmpty()) {
                continue;
            }
            // 为方法添加耗时统计
            String methodName = method.getName();
            String enhancedBody = buildEnhancedBody(method, methodName);
            // 替换方法体
            method.setBody(enhancedBody);
            System.out.println("增强方法: " + methodName);
        }
        // 返回增强后的类
        return ctClass.toClass();
    }
    /**
     * 构建增强的方法体
     */
    private static String buildEnhancedBody(CtMethod method, String methodName) throws Exception {
        StringBuilder sb = new StringBuilder();
        sb.append("{");
        sb.append("    long startTime = System.currentTimeMillis();");
        sb.append("    System.out.println(\"[Javassist] 调用方法: " + methodName + "\");");
        sb.append("    try {");
        // 如果是void方法
        if (method.getReturnType().equals(CtPrimitiveType.voidType)) {
            sb.append("        " + methodName + "$impl($$);");
        } else {
            sb.append("        Object result = " + methodName + "$impl($$);");
            sb.append("        return ($r) result;");
        }
        sb.append("    } finally {");
        sb.append("        long endTime = System.currentTimeMillis();");
        sb.append("        System.out.println(\"[Javassist] 方法: " + methodName + "耗时: \" + (endTime - startTime) + \"ms\");");
        sb.append("    }");
        sb.append("}");
        return sb.toString();
    }
    /**
     * 运行时增强(自定义ClassLoader方式)
     */
    public static Object createEnhancedInstance() throws Exception {
        ClassPool pool = ClassPool.getDefault();
        CtClass ctClass = pool.get("UserService");
        // 复制原有方法
        CtMethod method = ctClass.getDeclaredMethod("getUserInfo");
        // 创建增强方法
        String newMethodName = "getUserInfoEnhanced";
        CtMethod enhancedMethod = CtNewMethod.copy(method, newMethodName, ctClass, null);
        // 修改原方法
        method.setBody("{"
            + "long start = System.currentTimeMillis();"
            + "Object result = " + newMethodName + "($$);"
            + "long end = System.currentTimeMillis();"
            + "System.out.println(\"getUserInfo 耗时: \" + (end - start) + \"ms\");"
            + "return ($r) result;"
            + "}");
        // 添加增强方法
        ctClass.addMethod(enhancedMethod);
        // 创建实例
        Class<?> enhancedClass = ctClass.toClass();
        return enhancedClass.getDeclaredConstructor().newInstance();
    }
}

ASM实现字节码增强

import org.objectweb.asm.*;
import org.objectweb.asm.commons.AdviceAdapter;
import java.lang.reflect.Method;
/**
 * ASM字节码增强示例
 */
public class ASMEnhancer {
    /**
     * 自定义ClassVisitor
     */
    public static class EnhanceClassVisitor extends ClassVisitor {
        public EnhanceClassVisitor(ClassVisitor cv) {
            super(Opcodes.ASM9, cv);
        }
        @Override
        public MethodVisitor visitMethod(int access, String name, String descriptor, 
                                        String signature, String[] exceptions) {
            MethodVisitor mv = super.visitMethod(access, name, descriptor, signature, exceptions);
            // 只增强非抽象方法
            if (mv != null && !name.equals("<init>") && !name.equals("<clinit>")) {
                return new EnhanceMethodVisitor(mv, access, name, descriptor);
            }
            return mv;
        }
    }
    /**
     * 自定义MethodVisitor
     */
    public static class EnhanceMethodVisitor extends AdviceAdapter {
        private final String methodName;
        protected EnhanceMethodVisitor(MethodVisitor mv, int access, String name, String descriptor) {
            super(Opcodes.ASM9, mv, access, name, descriptor);
            this.methodName = name;
        }
        @Override
        protected void onMethodEnter() {
            // 方法进入时打印日志
            mv.visitFieldInsn(Opcodes.GETSTATIC, "java/lang/System", "out", "Ljava/io/PrintStream;");
            mv.visitLdcInsn("[ASM] 进入方法: " + methodName);
            mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL, "java/io/PrintStream", "println", "(Ljava/lang/String;)V", false);
            // 开始计时
            mv.visitMethodInsn(Opcodes.INVOKESTATIC, "java/lang/System", "currentTimeMillis", "()J", false);
            mv.visitVarInsn(Opcodes.LSTORE, 10);
        }
        @Override
        protected void onMethodExit(int opcode) {
            if (opcode != Opcodes.ATHROW) {
                // 计算耗时
                mv.visitMethodInsn(Opcodes.INVOKESTATIC, "java/lang/System", "currentTimeMillis", "()J", false);
                mv.visitVarInsn(Opcodes.LLOAD, 10);
                mv.visitInsn(Opcodes.LSUB);
                mv.visitVarInsn(Opcodes.LSTORE, 12);
                // 打印耗时
                mv.visitFieldInsn(Opcodes.GETSTATIC, "java/lang/System", "out", "Ljava/io/PrintStream;");
                mv.visitTypeInsn(Opcodes.NEW, "java/lang/StringBuilder");
                mv.visitInsn(Opcodes.DUP);
                mv.visitMethodInsn(Opcodes.INVOKESPECIAL, "java/lang/StringBuilder", "<init>", "()V", false);
                mv.visitLdcInsn("[ASM] 方法 " + methodName + " 耗时: ");
                mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL, "java/lang/StringBuilder", "append", "(Ljava/lang/String;)Ljava/lang/StringBuilder;", false);
                mv.visitVarInsn(Opcodes.LLOAD, 12);
                mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL, "java/lang/StringBuilder", "append", "(J)Ljava/lang/StringBuilder;", false);
                mv.visitLdcInsn("ms");
                mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL, "java/lang/StringBuilder", "append", "(Ljava/lang/String;)Ljava/lang/StringBuilder;", false);
                mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL, "java/lang/StringBuilder", "toString", "()Ljava/lang/String;", false);
                mv.visitMethodInsn(Opcodes.INVOKEVIRTUAL, "java/io/PrintStream", "println", "(Ljava/lang/String;)V", false);
            }
        }
    }
    /**
     * 使用ASM增强类
     */
    public static byte[] enhanceClass(byte[] originalClass) {
        ClassReader classReader = new ClassReader(originalClass);
        ClassWriter classWriter = new ClassWriter(classReader, ClassWriter.COMPUTE_MAXS);
        EnhanceClassVisitor visitor = new EnhanceClassVisitor(classWriter);
        classReader.accept(visitor, 0);
        return classWriter.toByteArray();
    }
}

Byte Buddy实现字节码增强

import net.bytebuddy.ByteBuddy;
import net.bytebuddy.agent.ByteBuddyAgent;
import net.bytebuddy.asm.Advice;
import net.bytebuddy.dynamic.DynamicType;
import net.bytebuddy.dynamic.loading.ClassLoadingStrategy;
import net.bytebuddy.implementation.MethodDelegation;
import net.bytebuddy.implementation.bind.annotation.*;
import net.bytebuddy.matcher.ElementMatchers;
import java.lang.reflect.Method;
import java.util.concurrent.Callable;
/**
 * Byte Buddy字节码增强示例(推荐使用)
 */
public class ByteBuddyEnhancer {
    /**
     * 方式一:使用Advice进行方法增强
     */
    public static Class<?> enhanceWithAdvice() throws Exception {
        DynamicType.Unloaded<UserService> unloaded = new ByteBuddy()
            .subclass(UserService.class)
            .method(ElementMatchers.named("getUserInfo"))
            .intercept(Advice.to(PerformanceAdvice.class))
            .method(ElementMatchers.named("updateUser"))
            .intercept(Advice.to(LoggingAdvice.class));
        return unloaded.load(UserService.class.getClassLoader())
                      .getLoaded();
    }
    /**
     * Advice类:性能监控
     */
    public static class PerformanceAdvice {
        @Advice.OnMethodEnter
        public static long enter(@Advice.Origin Method method) {
            System.out.println("[Byte Buddy] 进入方法: " + method.getName());
            return System.nanoTime();
        }
        @Advice.OnMethodExit
        public static void exit(@Advice.Enter long startTime, 
                                @Advice.Origin Method method,
                                @Advice.Return Object result) {
            long duration = (System.nanoTime() - startTime) / 1_000_000;
            System.out.println("[Byte Buddy] 方法 " + method.getName() + 
                             " 耗时: " + duration + "ms,返回值: " + result);
        }
    }
    /**
     * Advice类:日志记录
     */
    public static class LoggingAdvice {
        @Advice.OnMethodEnter
        public static void enter(@Advice.AllArguments Object[] args,
                                 @Advice.Origin Method method) {
            System.out.println("[Byte Buddy] 调用: " + method.getName() + 
                             ",参数: " + java.util.Arrays.toString(args));
        }
        @Advice.OnMethodExit
        public static void exit(@Advice.Origin Method method) {
            System.out.println("[Byte Buddy] 方法执行完成: " + method.getName());
        }
    }
    /**
     * 方式二:使用MethodDelegation
     */
    public static Class<?> enhanceWithDelegation() throws Exception {
        DynamicType.Unloaded<UserService> unloaded = new ByteBuddy()
            .subclass(UserService.class)
            .method(ElementMatchers.any())
            .intercept(MethodDelegation.to(Interceptor.class));
        return unloaded.load(UserService.class.getClassLoader())
                      .getLoaded();
    }
    /**
     * 拦截器类
     */
    public static class Interceptor {
        @RuntimeType
        public static Object intercept(@This Object target,
                                       @Origin Method method,
                                       @AllArguments Object[] args,
                                       @SuperCall Callable<?> callable) throws Exception {
            long startTime = System.currentTimeMillis();
            System.out.println("[Byte Buddy] 调用方法: " + method.getName());
            try {
                Object result = callable.call();
                System.out.println("[Byte Buddy] 方法返回: " + result);
                return result;
            } finally {
                long duration = System.currentTimeMillis() - startTime;
                System.out.println("[Byte Buddy] 方法耗时: " + duration + "ms");
            }
        }
    }
}

完整测试示例

import java.lang.reflect.Method;
/**
 * 测试类
 */
public class BytecodeEnhancementTest {
    public static void main(String[] args) throws Exception {
        System.out.println("========== 1. Javassist增强测试 ==========");
        testJavassist();
        System.out.println("\n========== 2. ASM增强测试 ==========");
        testASM();
        System.out.println("\n========== 3. Byte Buddy增强测试 ==========");
        testByteBuddy();
    }
    /**
     * 测试Javassist
     */
    private static void testJavassist() throws Exception {
        Class<?> enhancedClass = JavassistEnhancer.enhanceUserService();
        Object instance = enhancedClass.getDeclaredConstructor().newInstance();
        Method method = enhancedClass.getMethod("getUserInfo", String.class);
        Object result = method.invoke(instance, "1001");
        System.out.println("Javassist增强结果: " + result);
    }
    /**
     * 测试ASM
     */
    private static void testASM() throws Exception {
        // 读取原始类字节码
        ClassLoader classLoader = UserService.class.getClassLoader();
        java.io.InputStream is = classLoader.getResourceAsStream(
            "UserService.class".replace('.', '/')
        );
        byte[] originalBytes = is.readAllBytes();
        // 增强字节码
        byte[] enhancedBytes = ASMEnhancer.enhanceClass(originalBytes);
        // 自定义类加载器加载增强后的类
        ClassLoader enhancedClassLoader = new ClassLoader() {
            @Override
            public Class<?> loadClass(String name) throws ClassNotFoundException {
                if (name.equals("UserService")) {
                    return defineClass(name, enhancedBytes, 0, enhancedBytes.length);
                }
                return super.loadClass(name);
            }
        };
        Class<?> enhancedClass = enhancedClassLoader.loadClass("UserService");
        Object instance = enhancedClass.getDeclaredConstructor().newInstance();
        Method method = enhancedClass.getMethod("getUserInfo", String.class);
        Object result = method.invoke(instance, "1002");
        System.out.println("ASM增强结果: " + result);
    }
    /**
     * 测试Byte Buddy
     */
    private static void testByteBuddy() throws Exception {
        Class<?> enhancedClass = ByteBuddyEnhancer.enhanceWithAdvice();
        Object instance = enhancedClass.getDeclaredConstructor().newInstance();
        Method method = enhancedClass.getMethod("getUserInfo", String.class);
        Object result = method.invoke(instance, "1003");
        System.out.println("Byte Buddy增强结果: " + result);
        // 测试另一个方法
        Method updateMethod = enhancedClass.getMethod("updateUser", String.class, String.class);
        updateMethod.invoke(instance, "1003", "张三");
    }
}

实际应用案例:AOP日志框架

import java.lang.annotation.*;
import java.lang.reflect.Method;
/**
 * 自定义注解
 */
@Target(ElementType.METHOD)
@Retention(RetentionPolicy.RUNTIME)
public @interface Loggable {
    String value() default "";
}
/**
 * AOP日志增强器(企业级应用)
 */
public class LoggingEnhancer {
    /**
     * 使用Byte Buddy实现全面的AOP功能
     */
    public static Object enhance(Object target) throws Exception {
        DynamicType.Unloaded<Object> unloaded = new ByteBuddy()
            .subclass(target.getClass())
            .method(ElementMatchers.isAnnotatedWith(Loggable.class))
            .intercept(MethodDelegation.to(AspectLoggerInterceptor.class))
            .make();
        Class<?> enhancedClass = unloaded.load(target.getClass().getClassLoader())
                                       .getLoaded();
        return enhancedClass.getDeclaredConstructor().newInstance();
    }
    /**
     * 切面拦截器
     */
    public static class AspectLoggerInterceptor {
        @RuntimeType
        public static Object intercept(@This Object target,
                                       @Origin Method method,
                                       @AllArguments Object[] args,
                                       @SuperCall Callable<?> callable) throws Exception {
            // 前置通知
            LogContext context = new LogContext();
            context.setMethod(method);
            context.setArgs(args);
            context.setStartTime(System.currentTimeMillis());
            try {
                // 执行目标方法
                Object result = callable.call();
                // 后置通知
                logSuccess(context, result);
                return result;
            } catch (Exception e) {
                // 异常通知
                logError(context, e);
                throw e;
            } finally {
                // 最终通知
                logFinally(context);
            }
        }
        private static void logSuccess(LogContext ctx, Object result) {
            System.out.println(String.format(
                "[AOP] 方法 %s 成功执行,参数: %s,结果: %s",
                ctx.getMethod().getName(),
                Arrays.toString(ctx.getArgs()),
                result
            ));
        }
        private static void logError(LogContext ctx, Exception e) {
            System.err.println(String.format(
                "[AOP] 方法 %s 执行异常: %s",
                ctx.getMethod().getName(),
                e.getMessage()
            ));
        }
        private static void logFinally(LogContext ctx) {
            long duration = System.currentTimeMillis() - ctx.getStartTime();
            System.out.println(String.format(
                "[AOP] 方法 %s 执行耗时: %dms",
                ctx.getMethod().getName(),
                duration
            ));
        }
    }
    /**
     * 日志上下文
     */
    public static class LogContext {
        private Method method;
        private Object[] args;
        private long startTime;
        // getter/setter...
    }
}

使用注意事项

/**
 * 性能优化建议
 */
public class EnhancementBestPractices {
    /**
     * 1. 缓存增强类
     */
    public class ClassCache {
        private final Map<String, Class<?>> cache = new ConcurrentHashMap<>();
        public Class<?> getEnhancedClass(String className) {
            return cache.computeIfAbsent(className, this::createClass);
        }
        private Class<?> createClass(String className) {
            // 增强逻辑
            return UserService.class;
        }
    }
    /**
     * 2. 用于生产环境的Agent方式
     */
    public static class ByteBuddyAgentExample {
        public static void premain(String args, Instrumentation inst) {
            new AgentBuilder.Default()
                .type(ElementMatchers.nameStartsWith("com.example"))
                .transform((builder, typeDescription, classLoader, module) -> 
                    builder.method(ElementMatchers.any())
                           .intercept(MethodDelegation.to(Interceptor.class))
                )
                .installOn(inst);
        }
    }
}

字节码增强的核心价值在于:

  1. 非侵入式:不需要修改原有业务代码
  2. 性能监控:方法耗时、参数记录
  3. 安全控制:权限验证、数据脱敏
  4. 热部署:动态更新类定义

推荐优先使用 Byte Buddy(API简单、功能强大),ASM(性能最优但复杂度高),Javassist(简单但性能略差)。

抱歉,评论功能暂时关闭!