diff --git a/jero-boot/jero-boot-base/jero-boot-base-core/src/main/java/com/jero/config/CsrfFilter.java b/jero-boot/jero-boot-base/jero-boot-base-core/src/main/java/com/jero/config/CsrfFilter.java new file mode 100644 index 00000000..584801db --- /dev/null +++ b/jero-boot/jero-boot-base/jero-boot-base-core/src/main/java/com/jero/config/CsrfFilter.java @@ -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 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() { + + } +}