fix(请求拦截):正确处理文件类型的请求拦截

This commit is contained in:
wanjia 2024-06-03 15:36:36 +08:00
parent f00c095afa
commit ebbd2c0440
3 changed files with 50 additions and 59 deletions

View File

@ -77,6 +77,7 @@ public class ExceptionInterceptor {
@ResponseStatus(HttpStatus.INTERNAL_SERVER_ERROR)
@ExceptionHandler({Exception.class})
public ResponseData handle(HttpServletRequest request, Exception e) {
e.printStackTrace();
return DataPacketUtil.errorJsonResult(ResponseEnum.SYS_ERROR.getValue());
}

View File

@ -1,6 +1,8 @@
package com.xjy.ai.api.servlet;
import com.xjy.ai.api.wrapper.RequestWrapper;
import org.springframework.web.multipart.MultipartResolver;
import org.springframework.web.multipart.support.StandardServletMultipartResolver;
import org.springframework.web.servlet.DispatcherServlet;
import javax.servlet.http.HttpServletRequest;
@ -8,6 +10,8 @@ import javax.servlet.http.HttpServletResponse;
public class XinDispatcherServlet extends DispatcherServlet {
private final MultipartResolver multipartResolver = new StandardServletMultipartResolver();
/**
* 包装成我们自定义的request
* @param request
@ -16,6 +20,14 @@ public class XinDispatcherServlet extends DispatcherServlet {
*/
@Override
protected void doDispatch(HttpServletRequest request, HttpServletResponse response) throws Exception {
super.doDispatch(new RequestWrapper(request), response);
HttpServletRequest processedRequest = request;
if (multipartResolver.isMultipart(request)) {
processedRequest = multipartResolver.resolveMultipart(request);
} else {
processedRequest = new RequestWrapper(request);
}
super.doDispatch(processedRequest, response);
}
}

View File

@ -5,96 +5,74 @@ import javax.servlet.ServletInputStream;
import javax.servlet.ServletRequest;
import javax.servlet.http.HttpServletRequest;
import javax.servlet.http.HttpServletRequestWrapper;
import java.io.*;
import java.io.BufferedReader;
import java.io.ByteArrayInputStream;
import java.io.IOException;
import java.io.InputStreamReader;
import java.nio.charset.Charset;
import java.util.stream.Collectors;
/**
* 将request请求里面的inputStream复制一份出来放到固定参数里面
*
*/
public class RequestWrapper extends HttpServletRequestWrapper {
private final byte[] body;
public RequestWrapper(HttpServletRequest request) {
super(request);
String bodyStr = getBodyString(request);
body = bodyStr.getBytes(Charset.defaultCharset());
if (!isMultipartContent(request)) {
String bodyStr = getBodyString(request);
body = bodyStr.getBytes(Charset.defaultCharset());
} else {
body = new byte[0];
}
}
private boolean isMultipartContent(HttpServletRequest request) {
return request.getContentType() != null && request.getContentType().toLowerCase().startsWith("multipart/");
}
public String getBodyString(final ServletRequest request) {
try{
return inputStream2String(request.getInputStream());
}catch (IOException e){
throw new RuntimeException(e);
}
}
public String getBodyString() {
final InputStream inputStream = new ByteArrayInputStream(body);
return inputStream2String(inputStream);
}
/**
* 将inputStream里的数据读取出来并转换成字符串
*
* @param inputStream inputStream
* @return String
*/
private String inputStream2String(InputStream inputStream) {
StringBuilder sb = new StringBuilder();
BufferedReader reader = null;
try {
reader = new BufferedReader(new InputStreamReader(inputStream, Charset.defaultCharset()));
String line;
while ((line = reader.readLine()) != null) {
sb.append(line);
}
return inputStream2String(request.getInputStream());
} catch (IOException e) {
throw new RuntimeException(e);
} finally {
if (reader != null) {
try {
reader.close();
} catch (IOException e) {
throw new RuntimeException(e);
}
}
}
return sb.toString();
}
@Override
public BufferedReader getReader() throws IOException {
return new BufferedReader(new InputStreamReader(getInputStream()));
private String inputStream2String(ServletInputStream inputStream) {
return new BufferedReader(new InputStreamReader(inputStream, Charset.defaultCharset()))
.lines()
.collect(Collectors.joining("\n"));
}
@Override
public ServletInputStream getInputStream() throws IOException {
final ByteArrayInputStream inputStream = new ByteArrayInputStream(body);
return new ServletInputStream() {
@Override
public int read() throws IOException {
return inputStream.read();
}
private final ByteArrayInputStream byteArrayInputStream = new ByteArrayInputStream(body);
@Override
public boolean isFinished() {
return false;
return byteArrayInputStream.available() == 0;
}
@Override
public boolean isReady() {
return false;
return true;
}
@Override
public void setReadListener(ReadListener readListener) {
// No-op
}
@Override
public int read() throws IOException {
return byteArrayInputStream.read();
}
};
}
}
@Override
public BufferedReader getReader() throws IOException {
return new BufferedReader(new InputStreamReader(this.getInputStream()));
}
}