Java异步线程上下文案例

wen java案例 1

本文目录导读:

Java异步线程上下文案例

  1. 基础案例:ThreadLocal在异步场景的问题
  2. 解决方案:自定义ThreadLocal包装器
  3. 完整的异步上下文传递方案
  4. 使用现成库(TransmittableThreadLocal)
  5. 实际业务应用案例

我将为您提供一个Java异步线程上下文的完整案例,涵盖常见的场景和解决方案。

基础案例:ThreadLocal在异步场景的问题

import java.util.concurrent.*;
import java.util.HashMap;
import java.util.Map;
/**
 * 演示ThreadLocal在异步场景中的丢失问题
 */
public class AsyncThreadLocalDemo {
    // 模拟用户上下文
    private static ThreadLocal<String> userContext = new ThreadLocal<>();
    private static ThreadLocal<Map<String, Object>> requestContext = new ThreadLocal<>();
    public static void main(String[] args) throws Exception {
        ExecutorService executor = Executors.newFixedThreadPool(5);
        // 在主线程设置上下文
        userContext.set("User-001");
        Map<String, Object> attrs = new HashMap<>();
        attrs.put("requestId", "REQ-123");
        attrs.put("traceId", "TRACE-456");
        requestContext.set(attrs);
        System.out.println("主线程设置: " + userContext.get());
        System.out.println("主线程请求上下文: " + requestContext.get());
        // 使用submit提交任务
        Future<String> future = executor.submit(() -> {
            // 子线程中无法获取主线程的ThreadLocal
            return "子线程获取: " + userContext.get();
        });
        System.out.println(future.get());
        // 使用execute提交任务
        executor.execute(() -> {
            System.out.println("execute方式获取: " + userContext.get());
        });
        executor.shutdown();
    }
}

解决方案:自定义ThreadLocal包装器

import java.util.HashMap;
import java.util.Map;
import java.util.concurrent.*;
/**
 * 可传递的ThreadLocal包装器
 */
public class TransmittableThreadLocal<T> {
    // 存储当前线程的变量
    private static final ThreadLocal<Map<String, Object>> THREAD_LOCAL = new ThreadLocal<>();
    // 值提供者
    private final String key;
    private final T defaultValue;
    public TransmittableThreadLocal(String key, T defaultValue) {
        this.key = key;
        this.defaultValue = defaultValue;
    }
    public void set(T value) {
        Map<String, Object> map = getContextMap();
        map.put(key, value);
    }
    @SuppressWarnings("unchecked")
    public T get() {
        Map<String, Object> map = getContextMap();
        T value = (T) map.get(key);
        return value != null ? value : defaultValue;
    }
    public void remove() {
        Map<String, Object> map = getContextMap();
        map.remove(key);
    }
    private Map<String, Object> getContextMap() {
        Map<String, Object> map = THREAD_LOCAL.get();
        if (map == null) {
            map = new HashMap<>();
            THREAD_LOCAL.set(map);
        }
        return map;
    }
    /**
     * 获取整个上下文
     */
    public static Map<String, Object> capture() {
        Map<String, Object> map = THREAD_LOCAL.get();
        return map != null ? new ConcurrentHashMap<>(map) : null;
    }
    /**
     * 设置整个上下文
     */
    public static void replay(Map<String, Object> context) {
        if (context != null) {
            THREAD_LOCAL.set(new ConcurrentHashMap<>(context));
        }
    }
    /**
     * 清除上下文
     */
    public static void clear() {
        THREAD_LOCAL.remove();
    }
}

完整的异步上下文传递方案

import java.util.*;
import java.util.concurrent.*;
import java.util.function.Supplier;
/**
 * 完整的异步上下文传递解决方案
 */
public class AsyncContextDemo {
    private final ExecutorService executor;
    private final TransmittableThreadLocal<String> userId;
    private final TransmittableThreadLocal<String> traceId;
    private final TransmittableThreadLocal<Map<String, Object>> attributes;
    public AsyncContextDemo() {
        this.executor = Executors.newFixedThreadPool(10);
        this.userId = new TransmittableThreadLocal<>("userId", "unknown");
        this.traceId = new TransmittableThreadLocal<>("traceId", "no-trace");
        this.attributes = new TransmittableThreadLocal<>("attrs", new HashMap<>());
    }
    /**
     * 包装Runnable,传递上下文
     */
    public Runnable wrapRunnable(Runnable task) {
        // 捕获当前上下文
        Map<String, Object> context = TransmittableThreadLocal.capture();
        return () -> {
            try {
                // 在子线程中重放上下文
                TransmittableThreadLocal.replay(context);
                task.run();
            } finally {
                // 清理子线程上下文
                TransmittableThreadLocal.clear();
            }
        };
    }
    /**
     * 包装Callable,传递上下文
     */
    public <T> Callable<T> wrapCallable(Callable<T> task) {
        Map<String, Object> context = TransmittableThreadLocal.capture();
        return () -> {
            try {
                TransmittableThreadLocal.replay(context);
                return task.call();
            } finally {
                TransmittableThreadLocal.clear();
            }
        };
    }
    /**
     * 包装Supplier,传递上下文
     */
    public <T> Supplier<T> wrapSupplier(Supplier<T> supplier) {
        Map<String, Object> context = TransmittableThreadLocal.capture();
        return () -> {
            try {
                TransmittableThreadLocal.replay(context);
                return supplier.get();
            } finally {
                TransmittableThreadLocal.clear();
            }
        };
    }
    /**
     * 提交带上下文的任务
     */
    public Future<?> submit(Runnable task) {
        return executor.submit(wrapRunnable(task));
    }
    public <T> Future<T> submit(Callable<T> task) {
        return executor.submit(wrapCallable(task));
    }
    public void execute(Runnable task) {
        executor.execute(wrapRunnable(task));
    }
    /**
     * 测试示例
     */
    public void testDemo() throws Exception {
        // 设置上下文
        userId.set("User-100");
        traceId.set("TRACE-123456");
        attributes.get().put("requestId", "REQ-ABCDEF");
        System.out.println("主线程开始 - User: " + userId.get() 
            + ", Trace: " + traceId.get());
        // 示例1:使用submit执行
        Future<String> future = submit(() -> {
            // 模拟业务操作
            Thread.sleep(100);
            return "任务1 - User: " + userId.get() 
                + ", Trace: " + traceId.get() 
                + ", Attr: " + attributes.get().get("requestId");
        });
        System.out.println("子线程结果: " + future.get());
        // 示例2:使用execute执行
        execute(() -> {
            System.out.println("任务2 - User: " + userId.get());
        });
        // 示例3:使用CompletableFuture
        CompletableFuture<String> cf = CompletableFuture.supplyAsync(
            wrapSupplier(() -> "CompletableFuture - User: " + userId.get()),
            executor
        );
        System.out.println("CF结果: " + cf.get());
        // 示例4:异常处理
        Future<String> errorFuture = submit(() -> {
            if (true) throw new RuntimeException("模拟异常");
            return "永远执行不到";
        });
        try {
            errorFuture.get();
        } catch (Exception e) {
            System.out.println("捕获异常: " + e.getCause().getMessage());
        }
        executor.shutdown();
    }
    public static void main(String[] args) throws Exception {
        new AsyncContextDemo().testDemo();
    }
}

使用现成库(TransmittableThreadLocal)

import com.alibaba.ttl.TransmittableThreadLocal;
import com.alibaba.ttl.threadpool.TtlExecutors;
import java.util.concurrent.*;
/**
 * 使用阿里巴巴TTL库
 */
public class TTLDemo {
    // 使用TTL的ThreadLocal
    private static TransmittableThreadLocal<String> context = 
        new TransmittableThreadLocal<>();
    public static void main(String[] args) throws Exception {
        // 创建线程池
        ExecutorService executor = Executors.newFixedThreadPool(5);
        // 使用TtlExecutors包装线程池
        ExecutorService ttlExecutor = TtlExecutors.getTtlExecutorService(executor);
        // 主线程设置上下文
        context.set("TTL-Context-Value");
        System.out.println("主线程: " + context.get());
        // 提交任务
        ttlExecutor.execute(() -> {
            System.out.println("子线程获取: " + context.get());
        });
        // 使用FutureTask
        Future<String> future = ttlExecutor.submit(() -> {
            return "子线程返回值: " + context.get();
        });
        System.out.println(future.get());
        // 异步任务完成后再修改上下文
        Thread.sleep(1000);
        context.set("Modified-Value");
        ttlExecutor.execute(() -> {
            System.out.println("修改后的传递: " + context.get());
        });
        ttlExecutor.shutdown();
    }
}

实际业务应用案例

import java.util.*;
import java.util.concurrent.*;
/**
 * 业务场景:分布式追踪系统上下文传递
 */
public class BusinessContextDemo {
    // 业务上下文
    private static final ThreadLocal<TraceContext> TRACE_CONTEXT = new ThreadLocal<>();
    // 业务上下文类
    static class TraceContext {
        private final String traceId;
        private final String spanId;
        private final String userId;
        private final Map<String, String> tags = new HashMap<>();
        public TraceContext(String traceId, String spanId, String userId) {
            this.traceId = traceId;
            this.spanId = spanId;
            this.userId = userId;
        }
        @Override
        public String toString() {
            return String.format("TraceContext{traceId='%s', spanId='%s', userId='%s', tags=%s}",
                traceId, spanId, userId, tags);
        }
    }
    /**
     * 自定义线程池工厂,自动传递上下文
     */
    static class ContextAwareThreadFactory implements ThreadFactory {
        private final ThreadFactory delegate;
        public ContextAwareThreadFactory(ThreadFactory delegate) {
            this.delegate = delegate;
        }
        @Override
        public Thread newThread(Runnable r) {
            Thread thread = delegate.newThread(r);
            // 捕获当前请求上下文
            TraceContext context = TRACE_CONTEXT.get();
            if (context != null) {
                thread.setUncaughtExceptionHandler((t, e) -> {
                    System.err.println("线程异常: " + e.getMessage());
                });
            }
            return thread;
        }
    }
    /**
     * 异步任务管理器
     */
    static class AsyncTaskManager {
        private final ExecutorService executor;
        public AsyncTaskManager() {
            this.executor = createExecutor();
        }
        private ExecutorService createExecutor() {
            ThreadFactory factory = new ContextAwareThreadFactory(
                Executors.defaultThreadFactory()
            );
            return Executors.newFixedThreadPool(10, factory);
        }
        /**
         * 提交带上下文的异步任务
         */
        public <T> CompletableFuture<T> submitAsync(Callable<T> task) {
            TraceContext context = TRACE_CONTEXT.get();
            return CompletableFuture.supplyAsync(() -> {
                try {
                    // 在子线程恢复上下文
                    TraceContext oldContext = TRACE_CONTEXT.get();
                    TRACE_CONTEXT.set(context);
                    try {
                        return task.call();
                    } finally {
                        if (oldContext != null) {
                            TRACE_CONTEXT.set(oldContext);
                        } else {
                            TRACE_CONTEXT.remove();
                        }
                    }
                } catch (Exception e) {
                    throw new RuntimeException(e);
                }
            }, executor);
        }
        public void shutdown() {
            executor.shutdown();
        }
    }
    public static void main(String[] args) throws Exception {
        // 创建异步任务管理器
        AsyncTaskManager manager = new AsyncTaskManager();
        // 模拟一个请求
        String traceId = UUID.randomUUID().toString();
        String spanId = UUID.randomUUID().toString().substring(0, 8);
        String userId = "User-" + System.currentTimeMillis() % 1000;
        TraceContext context = new TraceContext(traceId, spanId, userId);
        context.tags.put("client", "web-app");
        context.tags.put("env", "production");
        // 设置请求上下文
        TRACE_CONTEXT.set(context);
        System.out.println("主线程开始处理请求: " + context);
        // 执行多个异步任务
        CompletableFuture<String> task1 = manager.submitAsync(() -> {
            Thread.sleep(100);
            TraceContext ctx = TRACE_CONTEXT.get();
            return "任务1完成 - " + ctx;
        });
        CompletableFuture<String> task2 = manager.submitAsync(() -> {
            Thread.sleep(200);
            TraceContext ctx = TRACE_CONTEXT.get();
            return "任务2完成 - " + ctx;
        });
        CompletableFuture<String> task3 = manager.submitAsync(() -> {
            Thread.sleep(150);
            TraceContext ctx = TRACE_CONTEXT.get();
            return "任务3完成 - " + ctx;
        });
        // 等待所有任务完成
        CompletableFuture.allOf(task1, task2, task3).join();
        System.out.println("任务1: " + task1.get());
        System.out.println("任务2: " + task2.get());
        System.out.println("任务3: " + task3.get());
        // 主线程上下文仍然保持
        System.out.println("主线程上下文: " + TRACE_CONTEXT.get());
        // 清理
        TRACE_CONTEXT.remove();
        manager.shutdown();
    }
}
  1. 问题ThreadLocal默认不会在线程间传递,子线程无法继承父线程的上下文。

  2. 解决方案

    • 手动通过构造函数传递
    • 使用包装器模式包装Runnable/Callable
    • 使用阿里巴巴的TransmittableThreadLocal
    • 自定义线程池工厂
  3. 最佳实践

    • 在异步任务入口捕获上下文
    • 在子线程中恢复上下文
    • 确保finally中清理上下文
    • 注意线程池复用导致的上下文混淆
  4. 性能考虑

    • 避免传递大对象
    • 使用合适的并发容器
    • 考虑内存泄漏风险

这些方案可以有效解决Java异步编程中的上下文传递问题,适用于分布式追踪、用户认证、请求日志等场景。

抱歉,评论功能暂时关闭!