【add】添加跨站过滤器和白名单txt文件用于读取白名单

This commit is contained in:
mzc5649
2021-04-16 15:50:51 +08:00
parent 5c0b28d4bf
commit 1e582930fa
@@ -0,0 +1,164 @@
package com.jero.config;
import cn.hutool.json.JSONObject;
import cn.hutool.json.JSONUtil;
import com.jero.common.api.vo.Result;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import org.springframework.util.ResourceUtils;
import javax.servlet.*;
import javax.servlet.http.HttpServletRequest;
import javax.servlet.http.HttpServletResponse;
import java.io.*;
import java.text.SimpleDateFormat;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Date;
import java.util.List;
/**
* 描述:跨站过滤器
*
* @Author: 马志朝
* @Date: 2021/4/16 14:09
*/
@Component
public class CsrfFilter implements Filter {
/**
* LOGGER
*/
private static final Log LOGGER = LogFactory.getLog(CsrfFilter.class);
/**
* 白名单
*/
private List<String> whiteUrls = new ArrayList<>();
/**
* size
*/
private int size = 0;
@Override
public void init(FilterConfig filterConfig) throws ServletException {
FileInputStream fis = null;
InputStreamReader isr = null;
BufferedReader br = null;
try {
File file = ResourceUtils.getFile("classpath:whiteUrls.txt");
fis = new FileInputStream(file);
isr = new InputStreamReader(fis);
br = new BufferedReader(isr);
String url = null;
while((url = br.readLine()) != null){
whiteUrls.add(url);
}
size = whiteUrls.size();
} catch (IOException e) {
e.printStackTrace();
} finally {
if(fis!=null){
try {
fis.close();
} catch (IOException e) {
e.printStackTrace();
}
}
if(isr!=null){
try {
isr.close();
} catch (IOException e) {
e.printStackTrace();
}
}
if(br!=null){
try {
br.close();
} catch (IOException e) {
e.printStackTrace();
}
}
}
}
@Override
public void doFilter(ServletRequest request, ServletResponse response, FilterChain chain) {
try {
HttpServletRequest req = (HttpServletRequest) request;
HttpServletResponse res = (HttpServletResponse) response;
// 获取请求url地址
String url = req.getRequestURL().toString();
// 获取来源
String referurl = req.getHeader("Referer");
if(isWhiteReq(referurl)){
chain.doFilter(request, response);
}else{
String log = "";
String date = new SimpleDateFormat("yyyyMMddHHmmss").format(new Date());
String ip = getIp(req);
log = "跨站请求---->>>" + ip + "||" + date + "||" + referurl + "||" + url;
LOGGER.warn(log);
//发送错误信息
JSONObject json = JSONUtil.parseObj(Result.error("监测到跨站请求,请求失败"), false);
res.setCharacterEncoding("UTF-8");
res.setContentType("application/json; charset=utf-8");
PrintWriter out = res.getWriter();
out.append(json.toStringPretty());
}
} catch (Exception e) {
LOGGER.error("doFilter", e);
}
}
/**
* 判断是否是白名单
*/
private boolean isWhiteReq(String referUrl) {
if (referUrl == null || "".equals(referUrl) || size == 0) {
return true;
} else {
String refHost = "";
referUrl = referUrl.toLowerCase();
if (referUrl.startsWith("http://")) {
refHost = referUrl.substring(7);
} else if (referUrl.startsWith("https://")) {
refHost = referUrl.substring(8);
}
for (String urlTemp : whiteUrls) {
if (refHost.contains(urlTemp.toLowerCase())) {
return true;
}
}
}
return false;
}
/**
* 获取登录用户IP地址
* @param request
* @return
*/
public String getIp(HttpServletRequest request) {
String ip = request.getHeader("x-forwarded-for");
if (ip == null || ip.length() == 0 || "unknown".equalsIgnoreCase(ip)) {
ip = request.getHeader("Proxy-Client-IP");
}
if (ip == null || ip.length() == 0 || "unknown".equalsIgnoreCase(ip)) {
ip = request.getHeader("WL-Proxy-Client-IP");
}
if (ip == null || ip.length() == 0 || "unknown".equalsIgnoreCase(ip)) {
ip = request.getRemoteAddr();
}
if (ip.equals("0:0:0:0:0:0:0:1")) {
ip = "localhost";
}
return ip;
}
@Override
public void destroy() {
}
}