内容简述
目的
解决servlet request InputStream body只能读取一次的问题
设置请求级别的线程变量 ThreadLocal
注意:Fitler 执行循序,以免在wrapper 初始化前,body已经消费
Request wrapper
wrapper
重写和 body 读取相关的 3 个核心方法:
- getInputStream()
- getReader()- getContentLength()
javaimport java.io.*; public class RepeatableHttpServletRequestWrapper extends HttpServletRequestWrapper { /** 缓存请求体字节数组 */ private byte[] body; /** * Constructs a request object wrapping the given request. * * @param request The request to wrap * @throws IllegalArgumentException if the request is null */ public RepeatableHttpServletRequestWrapper(HttpServletRequest request) throws IOException { super(request); this.body = toByteArray(request.getInputStream()); } /** * 将输入流读取为字节数组 * * @param input 输入流 * @return 字节数组 * @throws IOException 读取失败时抛出 */ private static byte[] toByteArray(InputStream input) throws IOException { ByteArrayOutputStream buffer = new ByteArrayOutputStream(); byte[] data = new byte[4096]; int n; while ((n = input.read(data, 0, data.length)) != -1) { buffer.write(data, 0, n); } return buffer.toByteArray(); } @Override public ServletInputStream getInputStream() { ByteArrayInputStream input = new ByteArrayInputStream(body); return new ServletInputStream() { @Override public boolean isFinished() { return false; } @Override public boolean isReady() { return false; } @Override public void setReadListener(ReadListener listener) { } @Override public int read(){ return input.read(); } }; } @Override public BufferedReader getReader(){ return new BufferedReader(new InputStreamReader(getInputStream())); } @Override public int getContentLength() { return body.length; } /** * 获取缓存的请求体字节数组 * * @return 请求体字节数组 */ public byte[] getBody() { return body; } }filter
importjakarta.servlet.FilterChain;importjakarta.servlet.ServletException;importjakarta.servlet.http.HttpServletRequest;importjakarta.servlet.http.HttpServletResponse;importorg.springframework.web.filter.OncePerRequestFilter;importjava.io.IOException;publicclassRepeatableReadFilterextendsOncePerRequestFilter{@OverrideprotectedvoiddoFilterInternal(HttpServletRequestrequest,HttpServletResponseresponse,FilterChainfilterChain)throwsServletException,IOException{if(request.getContentType()!=null){RepeatableHttpServletRequestWrapperwrapperRequest=newRepeatableHttpServletRequestWrapper(request);filterChain.doFilter(wrapperRequest,response);}else{filterChain.doFilter(request,response);}}}Request context holder
自定义Request context Object
importlombok.*;@Getter@Setter@NoArgsConstructor@AllArgsConstructor@ToStringpublicclassRequestContextHolderObject<T>{privateStringUUID;privateTt;}实现RequestContextHolder
publicclassRequestContextHolder{privatestaticfinalInheritableThreadLocal<Object>threadLocal=newInheritableThreadLocal<>();publicstatic<T>voidset(RequestContextHolderObject<T>commonHeader){threadLocal.set(commonHeader);}@SuppressWarnings("unchecked")publicstatic<T>RequestContextHolderObject<T>get(){return(RequestContextHolderObject<T>)threadLocal.get();}publicstaticvoidremove(){threadLocal.remove();}publicstaticStringgetUUID(){RequestContextHolderObject<Object>objectRequestContextHolderObject=get();if(objectRequestContextHolderObject==null){returnnull;}returnobjectRequestContextHolderObject.getUUID();}}Filter 实现
publicclassCustomerRequestContextFilterextendsOncePerRequestFilter{@OverrideprotectedvoiddoFilterInternal(HttpServletRequestrequest,HttpServletResponseresponse,FilterChainfilterChain)throwsServletException,IOException{// 继承 MDC 中有 traceId,没有生成新的try{StringtraceId=MDC.get("traceId");Stringuuid=StringUtil.isEmptyString(traceId)?UUID.randomUUID().toString():traceId;RequestContextHolder.set(newRequestContextHolderObject<Object>(uuid,"test"));filterChain.doFilter(request,response);}finally{RequestContextHolder.remove();}}}子线程或者线程池传递
- 修改RequestContextHolder 使用TTL
TransmittableThreadLocal (TTL) - 自定义线程池,添加传递Decorator 实现
importcom.ft.mall.common.web.requestholder.RequestContextHolder;importcom.ft.mall.common.web.requestholder.RequestContextHolderObject;importorg.slf4j.MDC;importorg.springframework.context.annotation.Bean;importorg.springframework.core.task.TaskDecorator;importorg.springframework.scheduling.concurrent.ThreadPoolTaskExecutor;importjava.util.Map;importjava.util.concurrent.Executor;publicclassSimpleExecutor{@Bean("asyncExecutor")publicExecutorasyncExecutor(){ThreadPoolTaskExecutorexecutor=newThreadPoolTaskExecutor();executor.setCorePoolSize(10);executor.setMaxPoolSize(20);executor.setQueueCapacity(100);executor.setThreadNamePrefix("async-");// 配置MDC上下文传递装饰器executor.setTaskDecorator(newMdcTaskDecorator());executor.initialize();returnexecutor;}/** * MDC上下文传递装饰器 */staticclassMdcTaskDecoratorimplementsTaskDecorator{@OverridepublicRunnabledecorate(Runnablerunnable){// 父线程拷贝MDC上下文Map<String,String>contextMap=MDC.getCopyOfContextMap();RequestContextHolderObject<Object>objectRequestContextHolderObject=RequestContextHolder.get();return()->{try{// 子线程设置MDCif(contextMap!=null){MDC.setContextMap(contextMap);RequestContextHolder.set(objectRequestContextHolderObject);}runnable.run();}finally{MDC.clear();RequestContextHolder.remove();}};}}}filter 注册
importcom.ft.mall.common.web.filter.RepeatableReadFilter;importjakarta.servlet.Filter;importorg.springframework.boot.web.servlet.FilterRegistrationBean;importorg.springframework.context.annotation.Bean;importorg.springframework.context.annotation.Configuration;importorg.springframework.core.Ordered;importorg.springframework.web.servlet.config.annotation.WebMvcConfigurer;@ConfigurationpublicclassWebConfigimplementsWebMvcConfigurer{/** * 注册可重复读请求体过滤器 * * <p>拦截所有路径,优先级最高,确保在 Spring 拦截器和其他 Filter 之前包装请求。 * * @return Filter 注册 Bean */@BeanpublicFilterRegistrationBean<Filter>repeatableReadFilter(){FilterRegistrationBean<Filter>registration=newFilterRegistrationBean<>();registration.setFilter(newRepeatableReadFilter());registration.addUrlPatterns("/*");registration.setOrder(Ordered.HIGHEST_PRECEDENCE);registration.setName("repeatbleFilter");returnregistration;}@BeanpublicFilterRegistrationBean<Filter>requestContextHolderFilter(){FilterRegistrationBean<Filter>registration=newFilterRegistrationBean<>();registration.setFilter(newCustomerRequestContextFilter());registration.addUrlPatterns("/*");registration.setOrder(Ordered.HIGHEST_PRECEDENCE+2);registration.setName("requestContextHolderFilter");returnregistration;}}