diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessor.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessor.java index 904455a41..ab61452d2 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessor.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessor.java @@ -20,11 +20,15 @@ import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; import java.lang.reflect.Modifier; import java.util.concurrent.Executor; +import java.util.concurrent.ExecutorService; +import java.util.function.Supplier; +import org.aopalliance.aop.Advice; import org.aopalliance.intercept.MethodInterceptor; import org.aopalliance.intercept.MethodInvocation; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; + import org.springframework.aop.framework.AopConfigException; import org.springframework.aop.framework.ProxyFactoryBean; import org.springframework.beans.BeansException; @@ -63,50 +67,97 @@ class ExecutorBeanPostProcessor implements BeanPostProcessor { @Override public Object postProcessAfterInitialization(Object bean, String beanName) throws BeansException { - if (bean instanceof Executor && !(bean instanceof ThreadPoolTaskExecutor)) { - Method execute = ReflectionUtils.findMethod(bean.getClass(), "execute", Runnable.class); - boolean methodFinal = Modifier.isFinal(execute.getModifiers()); - boolean classFinal = Modifier.isFinal(bean.getClass().getModifiers()); - boolean cglibProxy = !methodFinal && !classFinal; - Executor executor = (Executor) bean; - try { - return createProxy(bean, cglibProxy, executor); - } catch (AopConfigException e) { - if (cglibProxy) { - if (log.isDebugEnabled()) { - log.debug("Exception occurred while trying to create a proxy, falling back to JDK proxy", e); - } - return createProxy(bean, false, executor); - } - throw e; - } - } else if (bean instanceof ThreadPoolTaskExecutor) { - boolean classFinal = Modifier.isFinal(bean.getClass().getModifiers()); - boolean cglibProxy = !classFinal; - ThreadPoolTaskExecutor executor = (ThreadPoolTaskExecutor) bean; - return createThreadPoolTaskExecutorProxy(bean, cglibProxy, executor); + if (bean instanceof ThreadPoolTaskExecutor) { + return wrapThreadPoolTaskExecutor(bean); + } else if (bean instanceof ExecutorService) { + return wrapExecutorService(bean); + } else if (bean instanceof Executor) { + return wrapExecutor(bean); } return bean; } + private Object wrapExecutor(Object bean) { + Method execute = ReflectionUtils.findMethod(bean.getClass(), "execute", + Runnable.class); + boolean methodFinal = Modifier.isFinal(execute.getModifiers()); + boolean classFinal = Modifier.isFinal(bean.getClass().getModifiers()); + boolean cglibProxy = !methodFinal && !classFinal; + Executor executor = (Executor) bean; + try { + return createProxy(bean, cglibProxy, + new ExecutorMethodInterceptor(executor, this.beanFactory)); + } + catch (AopConfigException ex) { + if (cglibProxy) { + if (log.isDebugEnabled()) { + log.debug( + "Exception occurred while trying to create a proxy, falling back to JDK proxy", + ex); + } + return createProxy(bean, false, new ExecutorMethodInterceptor(executor, this.beanFactory)); + } + throw ex; + } + } + + private Object wrapThreadPoolTaskExecutor(Object bean) { + boolean classFinal = Modifier.isFinal(bean.getClass().getModifiers()); + boolean cglibProxy = !classFinal; + ThreadPoolTaskExecutor executor = (ThreadPoolTaskExecutor) bean; + return createThreadPoolTaskExecutorProxy(bean, cglibProxy, executor); + } + + private Object wrapExecutorService(Object bean) { + boolean classFinal = Modifier.isFinal(bean.getClass().getModifiers()); + boolean cglibProxy = !classFinal; + ExecutorService executor = (ExecutorService) bean; + return createExecutorServiceProxy(bean, cglibProxy, executor); + } + Object createThreadPoolTaskExecutorProxy(Object bean, boolean cglibProxy, ThreadPoolTaskExecutor executor) { + return getProxiedObject(bean, cglibProxy, executor, + () -> new LazyTraceThreadPoolTaskExecutor(this.beanFactory, executor)); + } + + Object createExecutorServiceProxy(Object bean, boolean cglibProxy, + ExecutorService executor) { + return getProxiedObject(bean, cglibProxy, executor, + () -> new TraceableExecutorService(this.beanFactory, executor)); + } + + private Object getProxiedObject(Object bean, boolean cglibProxy, Executor executor, + Supplier supplier) { ProxyFactoryBean factory = new ProxyFactoryBean(); factory.setProxyTargetClass(cglibProxy); - factory.addAdvice(new ExecutorMethodInterceptor(executor, this.beanFactory) { - @Override Executor executor(BeanFactory beanFactory, ThreadPoolTaskExecutor executor) { - return new LazyTraceThreadPoolTaskExecutor(beanFactory, executor); + factory.addAdvice(new ExecutorMethodInterceptor(executor, + this.beanFactory) { + @Override + T executor(BeanFactory beanFactory, T executor) { + return (T) supplier.get(); } }); factory.setTarget(bean); + try { + return getObject(factory); + } catch (Exception e) { + if (log.isDebugEnabled()) { + log.debug("Exception occurred while trying to get a proxy. Will fallback to a different implementation", e); + } + return supplier.get(); + } + } + + Object getObject(ProxyFactoryBean factory) { return factory.getObject(); } @SuppressWarnings("unchecked") - Object createProxy(Object bean, boolean cglibProxy, Executor executor) { + Object createProxy(Object bean, boolean cglibProxy, Advice advice) { ProxyFactoryBean factory = new ProxyFactoryBean(); factory.setProxyTargetClass(cglibProxy); - factory.addAdvice(new ExecutorMethodInterceptor(executor, this.beanFactory)); + factory.addAdvice(advice); factory.setTarget(bean); return factory.getObject(); } @@ -122,9 +173,9 @@ class ExecutorMethodInterceptor implements MethodInterceptor this.beanFactory = beanFactory; } - @Override public Object invoke(MethodInvocation invocation) - throws Throwable { - Executor executor = executor(this.beanFactory, this.delegate); + @Override + public Object invoke(MethodInvocation invocation) throws Throwable { + T executor = executor(this.beanFactory, this.delegate); Method methodOnTracedBean = getMethod(invocation, executor); if (methodOnTracedBean != null) { try { @@ -144,7 +195,7 @@ class ExecutorMethodInterceptor implements MethodInterceptor .findMethod(object.getClass(), method.getName(), method.getParameterTypes()); } - Executor executor(BeanFactory beanFactory, T executor) { - return new LazyTraceExecutor(beanFactory, executor); + T executor(BeanFactory beanFactory, T executor) { + return (T) new LazyTraceExecutor(beanFactory, executor); } } diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessorTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessorTests.java index 283c4ba73..d1b937d72 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessorTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/async/ExecutorBeanPostProcessorTests.java @@ -16,19 +16,38 @@ package org.springframework.cloud.sleuth.instrument.async; +import java.util.Collection; +import java.util.Collections; +import java.util.List; +import java.util.concurrent.Callable; +import java.util.concurrent.ExecutionException; import java.util.concurrent.Executor; +import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.Future; import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import brave.Tracer; import brave.Tracing; +import org.aopalliance.aop.Advice; +import org.junit.After; import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; +import org.mockito.BDDMockito; +import org.mockito.BDDMockito; import org.mockito.Mock; import org.mockito.Mockito; import org.mockito.junit.MockitoJUnitRunner; import org.springframework.aop.framework.AopConfigException; +import org.springframework.aop.framework.ProxyFactoryBean; +import org.springframework.aop.framework.ProxyFactoryBean; import org.springframework.beans.factory.BeanFactory; import org.springframework.cloud.sleuth.DefaultSpanNamer; import org.springframework.cloud.sleuth.SpanNamer; @@ -45,7 +64,23 @@ import static org.assertj.core.api.BDDAssertions.thenThrownBy; @RunWith(MockitoJUnitRunner.class) public class ExecutorBeanPostProcessorTests { - @Mock BeanFactory beanFactory; + @Mock + BeanFactory beanFactory; + Tracing tracing = Tracing.newBuilder().build(); + + + @Before + public void setup() { + Mockito.when(beanFactory.getBean(Tracing.class)) + .thenReturn(this.tracing); + Mockito.when(beanFactory.getBean(SpanNamer.class)) + .thenReturn(new DefaultSpanNamer()); + } + + @After + public void clear() { + this.tracing.close(); + } @Test public void should_create_a_cglib_proxy_by_default() throws Exception { @@ -63,30 +98,32 @@ public class ExecutorBeanPostProcessorTests { } @Test - public void should_create_jdk_proxy_when_cglib_fails_to_be_done() throws Exception { + public void should_fallback_to_sleuth_implementation_when_cglib_cannot_be_created() throws Exception { ScheduledExecutorService service = Executors.newSingleThreadScheduledExecutor(); Object o = new ExecutorBeanPostProcessor(this.beanFactory) .postProcessAfterInitialization(service, "foo"); - then(o).isInstanceOf(ScheduledExecutorService.class); - then(ClassUtils.isCglibProxy(o)).isFalse(); + then(o).isInstanceOf(TraceableExecutorService.class); service.shutdown(); } @Test - public void should_throw_exception_when_it_is_not_possible_to_create_any_proxy() throws Exception { + public void should_fallback_to_default_implementation_when_exception_thrown() + throws Exception { ScheduledExecutorService service = Executors.newSingleThreadScheduledExecutor(); ExecutorBeanPostProcessor bpp = new ExecutorBeanPostProcessor(this.beanFactory) { - @Override Object createProxy(Object bean, boolean cglibProxy, - Executor executor) { + + @Override + Object createProxy(Object bean, boolean cglibProxy, Advice advice) { throw new AopConfigException("foo"); } + }; - thenThrownBy(() -> bpp.postProcessAfterInitialization(service, "foo")) - .isInstanceOf(AopConfigException.class) - .hasMessage("foo"); + Object wrappedService = bpp.postProcessAfterInitialization(service, "foo"); + + then(wrappedService).isInstanceOf(TraceableExecutorService.class); service.shutdown(); } @@ -103,7 +140,8 @@ public class ExecutorBeanPostProcessorTests { } @Test - public void should_throw_exception_when_it_is_not_possible_to_create_any_proxyfor_ThreadPoolTaskExecutor() throws Exception { + public void should_throw_exception_when_it_is_not_possible_to_create_any_proxy_for_ThreadPoolTaskExecutor() + throws Exception { ThreadPoolTaskExecutor taskExecutor = new ThreadPoolTaskExecutor(); ExecutorBeanPostProcessor bpp = new ExecutorBeanPostProcessor(this.beanFactory) { @Override Object createThreadPoolTaskExecutorProxy(Object bean, boolean cglibProxy, @@ -113,8 +151,92 @@ public class ExecutorBeanPostProcessorTests { }; thenThrownBy(() -> bpp.postProcessAfterInitialization(taskExecutor, "foo")) - .isInstanceOf(AopConfigException.class) - .hasMessage("foo"); + .isInstanceOf(AopConfigException.class).hasMessage("foo"); + } + + @Test + public void should_fallback_to_sleuth_impl_when_it_is_not_possible_to_create_any_proxy_for_ExecutorService() + throws Exception { + ExecutorService service = BDDMockito.mock(ExecutorService.class); + ExecutorBeanPostProcessor bpp = new ExecutorBeanPostProcessor(this.beanFactory) { + @Override + Object getObject(ProxyFactoryBean factory) { + throw new AopConfigException("foo"); + } + }; + + Object o = bpp.postProcessAfterInitialization(service, "foo"); + + then(o).isInstanceOf(TraceableExecutorService.class); + } + + private ExecutorService exceptionThrowingExecutorService() { + return new ExecutorService() { + @Override + public void execute(Runnable command) { + + } + + @Override + public void shutdown() { + + } + + @Override + public List shutdownNow() { + return null; + } + + @Override + public boolean isShutdown() { + return false; + } + + @Override + public boolean isTerminated() { + return false; + } + + @Override + public boolean awaitTermination(long timeout, TimeUnit unit) throws InterruptedException { + return false; + } + + @Override + public Future submit(Callable task) { + throw new IllegalStateException("foo"); + } + + @Override + public Future submit(Runnable task, T result) { + return null; + } + + @Override + public Future submit(Runnable task) { + return null; + } + + @Override + public List> invokeAll(Collection> tasks) throws InterruptedException { + return null; + } + + @Override + public List> invokeAll(Collection> tasks, long timeout, TimeUnit unit) throws InterruptedException { + return null; + } + + @Override + public T invokeAny(Collection> tasks) throws InterruptedException, ExecutionException { + return null; + } + + @Override + public T invokeAny(Collection> tasks, long timeout, TimeUnit unit) throws InterruptedException, ExecutionException, TimeoutException { + return null; + } + }; } @Test